This commit is contained in:
Sourab Mangrulkar
2022-11-30 18:39:36 +05:30
parent 992422100f
commit a8350a57fe
2 changed files with 8 additions and 2 deletions
+7
View File
@@ -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:
+1 -2
View File
@@ -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():