From bb63dc34f48a56149247fcb972015fba84fb7c05 Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Mon, 5 Dec 2022 16:02:13 +0530 Subject: [PATCH] fixes --- src/pet/utils/save_and_load.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/pet/utils/save_and_load.py b/src/pet/utils/save_and_load.py index 1418dfa..ea2160d 100644 --- a/src/pet/utils/save_and_load.py +++ b/src/pet/utils/save_and_load.py @@ -12,17 +12,17 @@ def get_pet_model_state_dict(model): """ if model.pet_config.pet_type == PETType.LORA: - return lora_state_dict(model) + to_return = lora_state_dict(model, bias=model.pet_config.bias) else: to_return = {} state_dict = model.state_dict() 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(): - if any(module_name in key for module_name in model.modules_to_save): - to_return[key] = value - return to_return + if model.modules_to_save is not None: + for key, value in state_dict.items(): + if any(module_name in key for module_name in model.modules_to_save): + to_return[key] = value + return to_return def set_pet_model_state_dict(model, pet_model_state_dict):