From 122f708ae8551f5b88cfbe3b9579fcdb3cb1f7ec Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Tue, 4 Apr 2023 20:05:59 +0530 Subject: [PATCH] =?UTF-8?q?=F0=9F=98=85.=20Fix=20=F0=9F=90=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/peft/tuners/lora.py | 2 +- src/peft/utils/save_and_load.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 0023717..63201cc 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -404,7 +404,7 @@ class LoraLayer: nn.init.zeros_(self.lora_B[adapter_name].weight) -class Linear(nn.Linear): +class Linear(nn.Linear, LoraLayer): # Lora implemented in a dense layer def __init__( self, diff --git a/src/peft/utils/save_and_load.py b/src/peft/utils/save_and_load.py index 7ebdfaf..a258c32 100644 --- a/src/peft/utils/save_and_load.py +++ b/src/peft/utils/save_and_load.py @@ -53,7 +53,7 @@ def get_peft_model_state_dict(model, state_dict=None, adapter_name="default"): elif isinstance(config, PromptLearningConfig): to_return = {} if config.inference_mode: - prompt_embeddings = model.prompt_encoder.embedding.weight + prompt_embeddings = model.prompt_encoder[adapter_name].embedding.weight else: prompt_embeddings = model.get_prompt_embedding_to_save(adapter_name) to_return["prompt_embeddings"] = prompt_embeddings