mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
fixes
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user