diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 2db3e76..e2932d9 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -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) diff --git a/src/peft/tuners/prompt_tuning.py b/src/peft/tuners/prompt_tuning.py index 86f448c..1dead1d 100644 --- a/src/peft/tuners/prompt_tuning.py +++ b/src/peft/tuners/prompt_tuning.py @@ -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):