From 16ba5fcf035c84f50c4cc92d201111507e642c5c Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Fri, 2 Dec 2022 11:05:30 +0530 Subject: [PATCH] FSDP cpu offloading fix --- src/pet/pet_model.py | 30 +++++++++--------------------- 1 file changed, 9 insertions(+), 21 deletions(-) diff --git a/src/pet/pet_model.py b/src/pet/pet_model.py index 4ed76a1..8e96c04 100644 --- a/src/pet/pet_model.py +++ b/src/pet/pet_model.py @@ -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)