Merge pull request #117 from huggingface/smangrul/fix-lora-merging-in-inference

fix merging lora weights for inference
This commit is contained in:
Sourab Mangrulkar
2023-02-22 00:20:40 +05:30
committed by GitHub
2 changed files with 27 additions and 24 deletions
+7 -4
View File
@@ -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):
"""
+20 -20
View File
@@ -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: