mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
Merge pull request #150 from mayank31398/mayank/single-module
support option for encoder only prompts
This commit is contained in:
+16
-9
@@ -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):
|
||||
|
||||
@@ -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"})
|
||||
|
||||
Reference in New Issue
Block a user