mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
FSDP cpu offloading fix
This commit is contained in:
+9
-21
@@ -79,7 +79,7 @@ class PETModel(torch.nn.Module):
|
||||
Returns the prompt embedding to save when saving the model. Only applocable when `pet_config.pet_type !=
|
||||
PETType.LORA`.
|
||||
"""
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(1, -1).to(self.base_model.device)
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(1, -1).to("cuda")
|
||||
if self.pet_config.pet_type == PETType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : self.pet_config.num_virtual_tokens]
|
||||
prompt_embeddings = self.prompt_encoder(prompt_tokens)
|
||||
@@ -89,7 +89,7 @@ class PETModel(torch.nn.Module):
|
||||
"""
|
||||
Returns the virtual prompts to use for PET. Only applocable when `pet_config.pet_type != PETType.LORA`.
|
||||
"""
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(batch_size, -1).to(self.base_model.device)
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(batch_size, -1).to("cuda")
|
||||
if self.pet_config.pet_type == PETType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : self.pet_config.num_virtual_tokens]
|
||||
if self.pet_config.inference_mode:
|
||||
@@ -207,9 +207,7 @@ class PETModelForSequenceClassification(PETModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to("cuda")
|
||||
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.")
|
||||
@@ -230,7 +228,7 @@ class PETModelForSequenceClassification(PETModel):
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
kwargs["token_type_ids"] = torch.cat(
|
||||
(
|
||||
torch.zeros(batch_size, self.pet_config.num_virtual_tokens).to(self.base_model.device),
|
||||
torch.zeros(batch_size, self.pet_config.num_virtual_tokens).to("cuda"),
|
||||
kwargs["token_type_ids"],
|
||||
),
|
||||
dim=1,
|
||||
@@ -364,9 +362,7 @@ class PETModelForCausalLM(PETModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to("cuda")
|
||||
attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
@@ -393,9 +389,7 @@ class PETModelForCausalLM(PETModel):
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
# concat prompt labels
|
||||
if labels is not None:
|
||||
prefix_labels = torch.full((batch_size, self.pet_config.num_virtual_tokens), -100).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_labels = torch.full((batch_size, self.pet_config.num_virtual_tokens), -100).to("cuda")
|
||||
kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
@@ -459,9 +453,7 @@ class PETModelForSeq2SeqLM(PETModel):
|
||||
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.pet_config.num_virtual_tokens).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to("cuda")
|
||||
decoder_attention_mask = torch.cat((prefix_attention_mask, decoder_attention_mask), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
@@ -497,15 +489,11 @@ class PETModelForSeq2SeqLM(PETModel):
|
||||
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to("cuda")
|
||||
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.pet_config.num_virtual_tokens), -100).to(
|
||||
self.base_model.device
|
||||
)
|
||||
prefix_labels = torch.full((batch_size, self.pet_config.num_virtual_tokens), -100).to("cuda")
|
||||
kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
inputs_embeds = torch.cat((prompts[:, : self.pet_config.num_virtual_tokens], inputs_embeds), dim=1)
|
||||
|
||||
Reference in New Issue
Block a user