FSDP cpu offloading fix

This commit is contained in:
Sourab Mangrulkar
2022-12-02 11:05:30 +05:30
parent 7a9c4a6287
commit 16ba5fcf03
+9 -21
View File
@@ -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)