mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
fp32
This commit is contained in:
@@ -384,6 +384,7 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
prompts = prompts.to(inputs_embeds.dtype)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
@@ -542,6 +543,7 @@ class PeftModelForCausalLM(PeftModel):
|
||||
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, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
@@ -577,10 +579,10 @@ class PeftModelForCausalLM(PeftModel):
|
||||
model_kwargs["past_key_values"] = past_key_values
|
||||
else:
|
||||
if model_kwargs["past_key_values"] is None:
|
||||
inputs_embeds = self.word_embeddings(model_kwargs["input_ids"])
|
||||
prompts = self.get_prompt(batch_size=model_kwargs["input_ids"].shape[0])
|
||||
model_kwargs["inputs_embeds"] = torch.cat(
|
||||
(prompts, self.word_embeddings(model_kwargs["input_ids"])), dim=1
|
||||
)
|
||||
prompts = prompts.to(inputs_embeds.dtype)
|
||||
model_kwargs["inputs_embeds"] = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
model_kwargs["input_ids"] = None
|
||||
|
||||
return model_kwargs
|
||||
@@ -694,6 +696,7 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
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
|
||||
@@ -824,6 +827,7 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
prompts = prompts.to(inputs_embeds.dtype)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
|
||||
@@ -111,6 +111,7 @@ class PromptEmbedding(torch.nn.Module):
|
||||
init_token_ids = init_token_ids[:total_virtual_tokens]
|
||||
|
||||
word_embedding_weights = word_embeddings(torch.LongTensor(init_token_ids)).detach().clone()
|
||||
word_embedding_weights = word_embedding_weights.to(torch.float32)
|
||||
self.embedding.weight = torch.nn.Parameter(word_embedding_weights)
|
||||
|
||||
def forward(self, indices):
|
||||
|
||||
Reference in New Issue
Block a user