Merge pull request #106 from huggingface/smangrul/add-diable-adapter-context-manager

add disable adapter context manager
This commit is contained in:
Sourab Mangrulkar
2023-02-17 20:32:31 +05:30
committed by GitHub
2 changed files with 51 additions and 3 deletions
+17
View File
@@ -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):
"""
+34 -3
View File
@@ -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