From 173dc3dedf23650948e36c00892d431cf3751d77 Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Fri, 17 Feb 2023 17:40:45 +0530 Subject: [PATCH 1/2] add `disable_adapter` context manager --- src/peft/peft_model.py | 17 +++++++++++++++++ src/peft/tuners/lora.py | 33 +++++++++++++++++++++++++++++++-- 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index e2932d9..f12c5be 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -16,6 +16,7 @@ import inspect import os import warnings +from contextlib import contextmanager import torch from accelerate import dispatch_model, infer_auto_device_map @@ -288,6 +289,22 @@ class PeftModel(PushToHubMixin, torch.nn.Module): else: return self.base_model.model(*args, **kwargs) + @contextmanager + def disable_adapter(self): + """ + Disables the adapter module. + """ + if isinstance(self.peft_config, PromptLearningConfig): + old_forward = self.forward + self.forward = self.base_model.forward + else: + self.base_model.disable_adapter_layers() + yield + if isinstance(self.peft_config, PromptLearningConfig): + self.forward = old_forward + else: + self.base_model.enable_adapter_layers() + class PeftModelForSequenceClassification(PeftModel): """ diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 7808582..9b7e860 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -209,6 +209,17 @@ class LoraModel(torch.nn.Module): config["inference_mode"] = True return config + def _set_adapter_layers(self, enabled=True): + for module in self.model.modules(): + if isinstance(module, LoraLayer): + module.disable_adapters = False if enabled else True + + def enable_adapter_layers(self): + self._set_adapter_layers(enabled=True) + + def disable_adapter_layers(self): + self._set_adapter_layers(enabled=False) + # Below code is based on https://github.com/microsoft/LoRA/blob/main/loralib/layers.py # and modified to work with PyTorch FSDP @@ -257,6 +268,7 @@ class LoraLayer: # Mark the weight as unmerged self.merged = False self.merge_weights = merge_weights + self.disable_adapters = False class Linear(nn.Linear, LoraLayer): @@ -319,7 +331,14 @@ class Linear(nn.Linear, LoraLayer): self.merged = True def forward(self, x: torch.Tensor): - if self.r > 0 and not self.merged: + if self.disable_adapters: + if self.r > 0 and self.merged: + self.weight.data -= ( + transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling + ) + self.merged = False + return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + elif self.r > 0 and not self.merged: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) if self.r > 0: result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling @@ -413,7 +432,17 @@ class MergedLinear(nn.Linear, LoraLayer): self.merged = True def forward(self, x: torch.Tensor): - if self.merged: + if self.disable_adapters: + if self.r > 0 and self.merged 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 = False + return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + elif self.merged: return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) else: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) From 1a8928c5a405eb6c34909c890ec8da3f8db496ae Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Fri, 17 Feb 2023 17:48:16 +0530 Subject: [PATCH 2/2] Update lora.py --- src/peft/tuners/lora.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 9b7e860..82b8b64 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -494,6 +494,8 @@ if is_bnb_available(): def forward(self, x: torch.Tensor): result = super().forward(x) - if self.r > 0: + if self.disable_adapters: + return result + elif self.r > 0: result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling return result