From 2ab8b95867fe052df4e098b3f397ff7068cb1a90 Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Wed, 30 Nov 2022 19:10:05 +0530 Subject: [PATCH] fixes --- src/pet/tuners/lora.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/pet/tuners/lora.py b/src/pet/tuners/lora.py index 0036638..eb30bad 100644 --- a/src/pet/tuners/lora.py +++ b/src/pet/tuners/lora.py @@ -32,9 +32,9 @@ class LoRAModel(torch.nn.Module): def __init__(self, config, model): super().__init__() self.config = config - self.lora_model = model + self.model = model self.find_and_replace() - mark_only_lora_as_trainable(self.lora_model, self.config.bias) + mark_only_lora_as_trainable(self.model, self.config.bias) def find_and_replace(self): kwargs = { @@ -44,7 +44,7 @@ class LoRAModel(torch.nn.Module): "fan_in_fan_out": self.config.fan_in_fan_out, "merge_weights": self.config.merge_weights, } - key_list = [key for key, _ in self.lora_model.named_modules()] + key_list = [key for key, _ in self.model.named_modules()] for key in key_list: if any(key.endswith(target_key) for target_key in self.config.target_modules): parent, target, target_name = self.get_submodules(key) @@ -58,9 +58,9 @@ class LoRAModel(torch.nn.Module): self.replace_module(parent, target_name, new_module, target) def get_submodules(self, key): - parent = self.lora_model.get_submodule(".".join(key.split(".")[:-1])) + parent = self.model.get_submodule(".".join(key.split(".")[:-1])) target_name = key.split(".")[-1] - target = self.lora_model.get_submodule(key) + target = self.model.get_submodule(key) return parent, target, target_name def replace_module(self, parent_module, child_name, new_module, old_module): @@ -70,4 +70,4 @@ class LoRAModel(torch.nn.Module): new_module.bias = old_module.bias def forward(self, *args, **kwargs): - return self.lora_model(*args, **kwargs) + return self.model(*args, **kwargs)