diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 4df5cd9..b98fc9f 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -84,6 +84,7 @@ class PeftModel(PushToHubMixin, torch.nn.Module): self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.peft_config = {} self.active_adapter = adapter_name + self.peft_type = peft_config.peft_type if not isinstance(peft_config, PromptLearningConfig): self.peft_config[adapter_name] = peft_config self.base_model = PEFT_TYPE_TO_MODEL_MAPPING[peft_config.peft_type]( @@ -120,10 +121,10 @@ class PeftModel(PushToHubMixin, torch.nn.Module): if peft_config.base_model_name_or_path is None: peft_config.base_model_name_or_path = ( self.base_model.__dict__.get("name_or_path", None) - if isinstance(self.peft_config, PromptLearningConfig) + if isinstance(peft_config, PromptLearningConfig) else self.base_model.model.__dict__.get("name_or_path", None) ) - inference_mode = self.peft_config.inference_mode + inference_mode = peft_config.inference_mode peft_config.inference_mode = True peft_config.save_pretrained(output_dir) peft_config.inference_mode = inference_mode @@ -213,32 +214,33 @@ class PeftModel(PushToHubMixin, torch.nn.Module): """ Returns the virtual prompts to use for Peft. Only applicable when `peft_config.peft_type != PeftType.LORA`. """ + peft_config = self.active_peft_config prompt_encoder = self.prompt_encoder[self.active_adapter] prompt_tokens = self.prompt_tokens[self.active_adapter].unsqueeze(0).expand(batch_size, -1).to(self.device) - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: - prompt_tokens = prompt_tokens[:, : self.peft_config.num_virtual_tokens] - if self.peft_config.inference_mode: + if peft_config.peft_type == PeftType.PREFIX_TUNING: + prompt_tokens = prompt_tokens[:, : peft_config.num_virtual_tokens] + if peft_config.inference_mode: past_key_values = prompt_encoder.embedding.weight.repeat(batch_size, 1, 1) else: past_key_values = prompt_encoder(prompt_tokens) past_key_values = past_key_values.view( batch_size, - self.peft_config.num_virtual_tokens, - self.peft_config.num_layers * 2, - self.peft_config.num_attention_heads, - self.peft_config.token_dim // self.peft_config.num_attention_heads, + peft_config.num_virtual_tokens, + peft_config.num_layers * 2, + peft_config.num_attention_heads, + peft_config.token_dim // peft_config.num_attention_heads, ) - if self.peft_config.num_transformer_submodules == 2: + if peft_config.num_transformer_submodules == 2: past_key_values = torch.cat([past_key_values, past_key_values], dim=2) past_key_values = past_key_values.permute([2, 0, 3, 1, 4]).split( - self.peft_config.num_transformer_submodules * 2 + peft_config.num_transformer_submodules * 2 ) if TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING.get(self.config.model_type, None) is not None: post_process_fn = TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING[self.config.model_type] past_key_values = post_process_fn(past_key_values) return past_key_values else: - if self.peft_config.inference_mode: + if peft_config.inference_mode: prompts = prompt_encoder.embedding.weight.repeat(batch_size, 1, 1) else: prompts = prompt_encoder(prompt_tokens) @@ -281,13 +283,13 @@ class PeftModel(PushToHubMixin, torch.nn.Module): """ Disables the adapter module. """ - if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig): + if isinstance(self.active_peft_config, PromptLearningConfig): old_forward = self.forward self.forward = self.base_model.forward else: self.base_model.disable_adapter_layers() yield - if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig): + if isinstance(self.active_peft_config, PromptLearningConfig): self.forward = old_forward else: self.base_model.enable_adapter_layers() @@ -296,13 +298,14 @@ class PeftModel(PushToHubMixin, torch.nn.Module): """ Returns the base model. """ - return ( - self.base_model - if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig) - else self.base_model.model - ) + return self.base_model if isinstance(self.active_peft_config, PromptLearningConfig) else self.base_model.model def add_adapter(self, adapter_name, peft_config): + if peft_config.peft_type != self.peft_type: + raise ValueError( + f"Cannot combine adapters with different peft types. " + f"Found {self.peft_type} and {peft_config.peft_type}." + ) self.peft_config[adapter_name] = peft_config if isinstance(peft_config, PromptLearningConfig): self._setup_prompt_encoder(adapter_name) @@ -380,11 +383,9 @@ class PeftModel(PushToHubMixin, torch.nn.Module): **dispatch_model_kwargs, ) hook = AlignDevicesHook(io_same_device=True) - if not isinstance(self.peft_config[adapter_name]) == PeftType.LORA: - add_hook_to_module(self.base_model.model, hook) - else: + if isinstance(self.peft_config[adapter_name], PromptLearningConfig): remove_hook_from_submodules(self.prompt_encoder) - add_hook_to_module(self.base_model, hook) + add_hook_to_module(self.get_base_model(), hook) def set_adapter(self, adapter_name): """ @@ -397,6 +398,10 @@ class PeftModel(PushToHubMixin, torch.nn.Module): self.base_model.set_adapter(adapter_name) _set_adapter(self, adapter_name) + @property + def active_peft_config(self): + return self.peft_config[self.active_adapter] + class PeftModelForSequenceClassification(PeftModel): """ @@ -465,8 +470,8 @@ class PeftModelForSequenceClassification(PeftModel): **kwargs, ): return_dict = return_dict if return_dict is not None else self.config.use_return_dict - - if not isinstance(self.peft_config, PromptLearningConfig): + peft_config = self.active_peft_config + if not isinstance(peft_config, PromptLearningConfig): return self.base_model( input_ids=input_ids, attention_mask=attention_mask, @@ -481,7 +486,7 @@ class PeftModelForSequenceClassification(PeftModel): batch_size = input_ids.shape[0] if attention_mask is not None: # concat prompt attention mask - prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device) + prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device) attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1) if kwargs.get("position_ids", None) is not None: warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.") @@ -496,13 +501,13 @@ class PeftModelForSequenceClassification(PeftModel): } ) - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if peft_config.peft_type == PeftType.PREFIX_TUNING: return self._prefix_tuning_forward(input_ids=input_ids, **kwargs) else: if kwargs.get("token_type_ids", None) is not None: kwargs["token_type_ids"] = torch.cat( ( - torch.zeros(batch_size, self.peft_config.num_virtual_tokens).to(self.device), + torch.zeros(batch_size, peft_config.num_virtual_tokens).to(self.device), kwargs["token_type_ids"], ), dim=1, @@ -638,7 +643,8 @@ class PeftModelForCausalLM(PeftModel): return_dict=None, **kwargs, ): - if not isinstance(self.peft_config, PromptLearningConfig): + peft_config = self.active_peft_config + if not isinstance(peft_config, PromptLearningConfig): return self.base_model( input_ids=input_ids, attention_mask=attention_mask, @@ -653,7 +659,7 @@ class PeftModelForCausalLM(PeftModel): batch_size = input_ids.shape[0] if attention_mask is not None: # concat prompt attention mask - prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device) + prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device) attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1) if kwargs.get("position_ids", None) is not None: @@ -672,7 +678,7 @@ class PeftModelForCausalLM(PeftModel): } ) - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if peft_config.peft_type == PeftType.PREFIX_TUNING: past_key_values = self.get_prompt(batch_size) return self.base_model(input_ids=input_ids, past_key_values=past_key_values, **kwargs) else: @@ -680,7 +686,7 @@ class PeftModelForCausalLM(PeftModel): inputs_embeds = self.word_embeddings(input_ids) # concat prompt labels if labels is not None: - prefix_labels = torch.full((batch_size, self.peft_config.num_virtual_tokens), -100).to(self.device) + prefix_labels = torch.full((batch_size, peft_config.num_virtual_tokens), -100).to(self.device) kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1) prompts = self.get_prompt(batch_size=batch_size) prompts = prompts.to(inputs_embeds.dtype) @@ -688,9 +694,10 @@ class PeftModelForCausalLM(PeftModel): return self.base_model(inputs_embeds=inputs_embeds, **kwargs) def generate(self, **kwargs): + peft_config = self.active_peft_config self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation try: - if not isinstance(self.peft_config, PromptLearningConfig): + if not isinstance(peft_config, PromptLearningConfig): outputs = self.base_model.generate(**kwargs) else: if "input_ids" not in kwargs: @@ -698,13 +705,13 @@ class PeftModelForCausalLM(PeftModel): # For gpt2 models, we construct postion_ids on the fly by using attention mask, and position ids need to match input_shape. # for prefix tuning, input shape is determined using `input_ids`. Thus we should not expand 'attention_mask' here # for prompt tuning input_ids is not passed but a concatenated input_embeds is passed. Thus attention_mask needs to be of same size of num_virtual_tokens + input_ids - if kwargs.get("attention_mask", None) is not None and self.peft_config.peft_type in [ + if kwargs.get("attention_mask", None) is not None and peft_config.peft_type in [ PeftType.PROMPT_TUNING, PeftType.P_TUNING, ]: # concat prompt attention mask prefix_attention_mask = torch.ones( - kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens + kwargs["input_ids"].shape[0], peft_config.num_virtual_tokens ).to(kwargs["input_ids"].device) kwargs["attention_mask"] = torch.cat((prefix_attention_mask, kwargs["attention_mask"]), dim=1) @@ -728,17 +735,18 @@ class PeftModelForCausalLM(PeftModel): return outputs def prepare_inputs_for_generation(self, *args, **kwargs): + peft_config = self.active_peft_config model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs) - if isinstance(self.peft_config, PromptLearningConfig): - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if isinstance(peft_config, PromptLearningConfig): + if peft_config.peft_type == PeftType.PREFIX_TUNING: prefix_attention_mask = torch.ones( - model_kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens + model_kwargs["input_ids"].shape[0], peft_config.num_virtual_tokens ).to(model_kwargs["input_ids"].device) model_kwargs["attention_mask"] = torch.cat( (prefix_attention_mask, model_kwargs["attention_mask"]), dim=1 ) - if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if model_kwargs["past_key_values"] is None and peft_config.peft_type == PeftType.PREFIX_TUNING: past_key_values = self.get_prompt(batch_size=model_kwargs["input_ids"].shape[0]) model_kwargs["past_key_values"] = past_key_values else: @@ -810,7 +818,8 @@ class PeftModelForSeq2SeqLM(PeftModel): return_dict=None, **kwargs, ): - if not isinstance(self.peft_config, PromptLearningConfig): + peft_config = self.active_peft_config + if not isinstance(peft_config, PromptLearningConfig): return self.base_model( input_ids=input_ids, attention_mask=attention_mask, @@ -828,7 +837,7 @@ class PeftModelForSeq2SeqLM(PeftModel): batch_size = input_ids.shape[0] if decoder_attention_mask is not None: # concat prompt attention mask - prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device) + prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device) decoder_attention_mask = torch.cat((prefix_attention_mask, decoder_attention_mask), dim=1) if kwargs.get("position_ids", None) is not None: @@ -848,7 +857,7 @@ class PeftModelForSeq2SeqLM(PeftModel): } ) - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if peft_config.peft_type == PeftType.PREFIX_TUNING: past_key_values = self.get_prompt(batch_size) return self.base_model( input_ids=input_ids, decoder_input_ids=decoder_input_ids, past_key_values=past_key_values, **kwargs @@ -864,35 +873,36 @@ class PeftModelForSeq2SeqLM(PeftModel): if attention_mask is not None: # concat prompt attention mask - prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device) + prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device) kwargs["attention_mask"] = torch.cat((prefix_attention_mask, attention_mask), dim=1) # concat prompt labels if labels is not None: - if self.peft_config.num_transformer_submodules == 1: + if peft_config.num_transformer_submodules == 1: kwargs["labels"] = labels - elif self.peft_config.num_transformer_submodules == 2: - prefix_labels = torch.full((batch_size, self.peft_config.num_virtual_tokens), -100).to(self.device) + elif peft_config.num_transformer_submodules == 2: + prefix_labels = torch.full((batch_size, peft_config.num_virtual_tokens), -100).to(self.device) kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1) prompts = self.get_prompt(batch_size=batch_size) prompts = prompts.to(inputs_embeds.dtype) - inputs_embeds = torch.cat((prompts[:, : self.peft_config.num_virtual_tokens], inputs_embeds), dim=1) - if self.peft_config.num_transformer_submodules == 1: + inputs_embeds = torch.cat((prompts[:, : peft_config.num_virtual_tokens], inputs_embeds), dim=1) + if peft_config.num_transformer_submodules == 1: return self.base_model(inputs_embeds=inputs_embeds, **kwargs) - elif self.peft_config.num_transformer_submodules == 2: + elif peft_config.num_transformer_submodules == 2: decoder_inputs_embeds = torch.cat( - (prompts[:, self.peft_config.num_virtual_tokens :], decoder_inputs_embeds), dim=1 + (prompts[:, peft_config.num_virtual_tokens :], decoder_inputs_embeds), dim=1 ) return self.base_model( inputs_embeds=inputs_embeds, decoder_inputs_embeds=decoder_inputs_embeds, **kwargs ) def generate(self, **kwargs): + peft_config = self.active_peft_config self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation self.base_model._prepare_encoder_decoder_kwargs_for_generation = ( self._prepare_encoder_decoder_kwargs_for_generation ) try: - if not isinstance(self.peft_config, PromptLearningConfig): + if not isinstance(peft_config, PromptLearningConfig): outputs = self.base_model.generate(**kwargs) else: if "input_ids" not in kwargs: @@ -908,7 +918,7 @@ class PeftModelForSeq2SeqLM(PeftModel): ) kwargs["token_type_ids"] = None - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if peft_config.peft_type == PeftType.PREFIX_TUNING: outputs = self.base_model.generate(**kwargs) else: raise NotImplementedError @@ -926,8 +936,9 @@ class PeftModelForSeq2SeqLM(PeftModel): return outputs def prepare_inputs_for_generation(self, *args, **kwargs): + peft_config = self.active_peft_config model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs) - if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if model_kwargs["past_key_values"] is None and peft_config.peft_type == PeftType.PREFIX_TUNING: batch_size = model_kwargs["decoder_input_ids"].shape[0] past_key_values = self.get_prompt(batch_size) model_kwargs["past_key_values"] = past_key_values @@ -1000,9 +1011,10 @@ class PeftModelForTokenClassification(PeftModel): return_dict=None, **kwargs, ): + peft_config = self.active_peft_config return_dict = return_dict if return_dict is not None else self.config.use_return_dict - if not isinstance(self.peft_config, PromptLearningConfig): + if not isinstance(peft_config, PromptLearningConfig): return self.base_model( input_ids=input_ids, attention_mask=attention_mask, @@ -1017,7 +1029,7 @@ class PeftModelForTokenClassification(PeftModel): batch_size = input_ids.shape[0] if attention_mask is not None: # concat prompt attention mask - prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device) + prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device) attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1) if kwargs.get("position_ids", None) is not None: warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.") @@ -1032,13 +1044,13 @@ class PeftModelForTokenClassification(PeftModel): } ) - if self.peft_config.peft_type == PeftType.PREFIX_TUNING: + if peft_config.peft_type == PeftType.PREFIX_TUNING: return self._prefix_tuning_forward(input_ids=input_ids, **kwargs) else: if kwargs.get("token_type_ids", None) is not None: kwargs["token_type_ids"] = torch.cat( ( - torch.zeros(batch_size, self.peft_config.num_virtual_tokens).to(self.device), + torch.zeros(batch_size, peft_config.num_virtual_tokens).to(self.device), kwargs["token_type_ids"], ), dim=1, diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 977a3b7..0b96a77 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -323,17 +323,7 @@ class LoraModel(torch.nn.Module): if isinstance(target, LoraLayer): bias = target.bias is not None new_module = torch.nn.Linear(target.in_features, target.out_features, bias=bias) - - # manually merge if not merged - if not target.merged: - # merge weights per: https://arxiv.org/pdf/2106.09685.pdf / page 4 - if target.r > 0: - target.weight.data += ( - transpose(target.lora_B.weight @ target.lora_A.weight, target.fan_in_fan_out) - * target.scaling - ).to(target.weight.dtype) - target.merged = True - + target.merge() self._replace_module(parent, target_name, new_module, target) return self.model