diff --git a/src/pet/pet_model.py b/src/pet/pet_model.py index a3f6fa6..1398146 100644 --- a/src/pet/pet_model.py +++ b/src/pet/pet_model.py @@ -53,6 +53,13 @@ class PETModel(torch.nn.Module): self.pet_config.num_virtual_tokens * self.pet_config.num_transformer_submodules ).long() + def get_prompt_embedding_to_save(self): + prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(1, -1).to(self.base_model.device) + 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) + return prompt_embeddings[0].detach().cpu() + def get_prompt(self, batch_size): prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(batch_size, -1).to(self.base_model.device) if self.pet_config.pet_type == PETType.PREFIX_TUNING: diff --git a/src/pet/utils/save_and_load.py b/src/pet/utils/save_and_load.py index 39fa579..e2eea05 100644 --- a/src/pet/utils/save_and_load.py +++ b/src/pet/utils/save_and_load.py @@ -9,8 +9,7 @@ def get_pet_model_state_dict(model): else: to_return = {} state_dict = model.state_dict() - prompt_tokens = model.prompt_tokens.unsqueeze(0).expand(1, -1).to(model.base_model.device) - prompt_embeddings = model.prompt_encoder(prompt_tokens).detach().cpu() + prompt_embeddings = model.get_prompt_embedding_to_save() to_return["prompt_embeddings"] = prompt_embeddings if model.modules_to_save is not None: for key, value in state_dict.items():