From e6ef85a71198c1e49916ad0e20c8c0fa23b6ea7f Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Wed, 22 Feb 2023 00:00:36 +0530 Subject: [PATCH 1/2] fix merging lora weights for inference --- src/peft/tuners/lora.py | 40 ++++++++++++++++++++-------------------- 1 file changed, 20 insertions(+), 20 deletions(-) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 82b8b64..0a8e1ca 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -132,7 +132,7 @@ class LoraModel(torch.nn.Module): "lora_alpha": self.peft_config.lora_alpha, "lora_dropout": self.peft_config.lora_dropout, "fan_in_fan_out": self.peft_config.fan_in_fan_out, - "merge_weights": self.peft_config.merge_weights, + "merge_weights": self.peft_config.merge_weights or self.peft_config.inference_mode, } key_list = [key for key, _ in self.model.named_modules()] for key in key_list: @@ -310,7 +310,14 @@ class Linear(nn.Linear, LoraLayer): nn.Linear.train(self, mode) self.lora_A.train(mode) self.lora_B.train(mode) - if self.merge_weights and self.merged: + if not mode and self.merge_weights and not self.merged: + # Merge the weights and mark it + if self.r > 0: + self.weight.data += ( + transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling + ) + self.merged = True + elif self.merge_weights and self.merged: # Make sure that the weights are not merged if self.r > 0: self.weight.data -= ( @@ -322,13 +329,6 @@ class Linear(nn.Linear, LoraLayer): nn.Linear.eval(self) self.lora_A.eval() self.lora_B.eval() - if self.merge_weights and not self.merged: - # Merge the weights and mark it - if self.r > 0: - self.weight.data += ( - transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling - ) - self.merged = True def forward(self, x: torch.Tensor): if self.disable_adapters: @@ -405,7 +405,17 @@ class MergedLinear(nn.Linear, LoraLayer): nn.Linear.train(self, mode) self.lora_A.train(mode) self.lora_B.train(mode) - if self.merge_weights and self.merged: + if not mode and self.merge_weights and not self.merged: + # Merge the weights and mark it + if self.r > 0 and any(self.enable_lora): + delta_w = F.conv1d( + self.lora_A.weight.data.unsqueeze(0), + self.lora_B.weight.data.unsqueeze(-1), + groups=sum(self.enable_lora), + ).squeeze(0) + self.weight.data += self.zero_pad(transpose(delta_w * self.scaling, self.fan_in_fan_out)) + self.merged = True + elif self.merge_weights and self.merged: # Make sure that the weights are not merged if self.r > 0 and any(self.enable_lora): delta_w = F.conv1d( @@ -420,16 +430,6 @@ class MergedLinear(nn.Linear, LoraLayer): nn.Linear.eval(self) self.lora_A.eval() self.lora_B.eval() - if self.merge_weights and not self.merged: - # Merge the weights and mark it - if self.r > 0 and any(self.enable_lora): - delta_w = F.conv1d( - self.lora_A.weight.data.unsqueeze(0), - self.lora_B.weight.data.unsqueeze(-1), - groups=sum(self.enable_lora), - ).squeeze(0) - self.weight.data += self.zero_pad(transpose(delta_w * self.scaling, self.fan_in_fan_out)) - self.merged = True def forward(self, x: torch.Tensor): if self.disable_adapters: From 1ef0f89a0c0125eb189cb0932d8da78919ca17d5 Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Wed, 22 Feb 2023 00:14:24 +0530 Subject: [PATCH 2/2] add util for getting the base model --- src/peft/peft_model.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index f12c5be..e20e64d 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -284,10 +284,7 @@ class PeftModel(PushToHubMixin, torch.nn.Module): """ Forward pass of the model. """ - if isinstance(self.peft_config, PromptLearningConfig): - return self.base_model(*args, **kwargs) - else: - return self.base_model.model(*args, **kwargs) + return self.get_base_model()(*args, **kwargs) @contextmanager def disable_adapter(self): @@ -305,6 +302,12 @@ class PeftModel(PushToHubMixin, torch.nn.Module): else: self.base_model.enable_adapter_layers() + def get_base_model(self): + """ + Returns the base model. + """ + return self.base_model if isinstance(self.peft_config, PromptLearningConfig) else self.base_model.model + class PeftModelForSequenceClassification(PeftModel): """