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..82b8b64 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) @@ -465,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