diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 4703059..b0c9bac 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -184,7 +184,6 @@ class PeftModel(PushToHubMixin, torch.nn.Module): return model def _setup_prompt_encoder(self): - num_transformer_submodules = 0 transformer_backbone = None for name, module in self.base_model.named_children(): for param in module.parameters(): @@ -194,8 +193,9 @@ class PeftModel(PushToHubMixin, torch.nn.Module): if transformer_backbone is None: transformer_backbone = module self.transformer_backbone_name = name - num_transformer_submodules += 1 - self.peft_config.num_transformer_submodules = 2 if self.peft_config.task_type == TaskType.SEQ_2_SEQ_LM else 1 + + if self.peft_config.num_transformer_submodules is None: + self.peft_config.num_transformer_submodules = 2 if self.peft_config.task_type == TaskType.SEQ_2_SEQ_LM else 1 for named_param, value in list(transformer_backbone.named_parameters()): if value.shape[0] == self.base_model.config.vocab_size: @@ -719,15 +719,22 @@ class PeftModelForSeq2SeqLM(PeftModel): kwargs["attention_mask"] = torch.cat((prefix_attention_mask, attention_mask), dim=1) # concat prompt labels if labels is not None: - prefix_labels = torch.full((batch_size, self.peft_config.num_virtual_tokens), -100).to(self.device) - kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1) + if self.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) + 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) - decoder_inputs_embeds = torch.cat( - (prompts[:, self.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) + if self.peft_config.num_transformer_submodules == 1: + return self.base_model(inputs_embeds=inputs_embeds, **kwargs) + elif self.peft_config.num_transformer_submodules == 2: + decoder_inputs_embeds = torch.cat( + (prompts[:, self.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): if not isinstance(self.peft_config, PromptLearningConfig): diff --git a/src/peft/utils/config.py b/src/peft/utils/config.py index f0587fe..c0a55f5 100644 --- a/src/peft/utils/config.py +++ b/src/peft/utils/config.py @@ -161,6 +161,6 @@ class PromptLearningConfig(PeftConfig): token_dim: int = field( default=None, metadata={"help": "The hidden embedding dimension of the base transformer model"} ) - num_transformer_submodules: Optional[int] = field(default=1, metadata={"help": "Number of transformer submodules"}) + num_transformer_submodules: Optional[int] = field(default=None, metadata={"help": "Number of transformer submodules"}) num_attention_heads: Optional[int] = field(default=None, metadata={"help": "Number of attention heads"}) num_layers: Optional[int] = field(default=None, metadata={"help": "Number of transformer layers"})