fix half precision forward

This commit is contained in:
younesbelkada
2023-04-04 10:01:47 +00:00
parent dd30335ffd
commit 8266e2ee4f
3 changed files with 64 additions and 11 deletions
+17
View File
@@ -81,6 +81,7 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
self.modules_to_save = self.peft_config.modules_to_save
_set_trainable(self)
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.base_model_torch_dtype = getattr(model, "dtype", None)
def save_pretrained(self, save_directory, **kwargs):
r"""
@@ -673,6 +674,22 @@ class PeftModelForCausalLM(PeftModel):
if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING:
past_key_values = self.get_prompt(batch_size=model_kwargs["input_ids"].shape[0])
if self.base_model_torch_dtype is not None:
# handle the case for Bloom where it outputs tuple of tuples
if isinstance(past_key_values[0], tuple):
past_key_values = tuple(
tuple(
past_key_value.to(self.base_model_torch_dtype)
for past_key_value in past_key_value_tuple
)
for past_key_value_tuple in past_key_values
)
else:
past_key_values = tuple(
past_key_value.to(self.base_model_torch_dtype) for past_key_value in past_key_values
)
model_kwargs["past_key_values"] = past_key_values
else:
if model_kwargs["past_key_values"] is None:
+24 -11
View File
@@ -394,21 +394,26 @@ class Linear(nn.Linear, LoraLayer):
self.lora_B.eval()
def forward(self, x: torch.Tensor):
previous_dtype = self.weight.dtype
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
)
matmul_output = self.lora_B.weight @ self.lora_A.weight
self.weight.data -= transpose(matmul_output.to(previous_dtype), 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)
result = 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
return result
result += self.lora_B(self.lora_A(self.lora_dropout(x.to(self.lora_A.weight.dtype)))) * self.scaling
else:
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
if result.dtype != previous_dtype:
result = result.to(previous_dtype)
return result
class MergedLinear(nn.Linear, LoraLayer):
@@ -508,6 +513,8 @@ class MergedLinear(nn.Linear, LoraLayer):
self.lora_B.eval()
def forward(self, x: torch.Tensor):
previous_dtype = x.dtype
if self.disable_adapters:
if self.r > 0 and self.merged and any(self.enable_lora):
delta_w = (
@@ -519,18 +526,24 @@ class MergedLinear(nn.Linear, LoraLayer):
.squeeze(0)
.transpose(-2, -1)
)
delta_w = delta_w.to(self.weight.dtype)
self.weight.data -= transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
self.merged = False
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
result = 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)
result = 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)
if self.r > 0:
after_A = self.lora_A(self.lora_dropout(x))
after_A = self.lora_A(self.lora_dropout(x.to(self.lora_A.weight.dtype)))
after_B = self.lora_B(after_A.transpose(-2, -1)).transpose(-2, -1)
result += self.zero_pad(after_B) * self.scaling
return result
result = result.to(previous_dtype)
return result
if is_bnb_available():
+23
View File
@@ -236,3 +236,26 @@ class PeftModelTester(unittest.TestCase, PeftTestMixin):
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
def test_generate(self, test_name, model_id, config_cls, config_kwargs):
self._test_generate(model_id, config_cls, config_kwargs)
def _test_generate_half_prec(self, model_id, config_cls, config_kwargs):
model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype=torch.bfloat16)
config = config_cls(
base_model_name_or_path=model_id,
**config_kwargs,
)
model = get_peft_model(model, config)
model = model.to(self.torch_device)
input_ids = torch.LongTensor([[1, 1, 1], [2, 1, 2]]).to(self.torch_device)
attention_mask = torch.LongTensor([[1, 1, 1], [1, 0, 1]]).to(self.torch_device)
# check if `generate` works
_ = model.generate(input_ids=input_ids, attention_mask=attention_mask)
with self.assertRaises(TypeError):
# check if `generate` raises an error if no positional arguments are passed
_ = model.generate(input_ids, attention_mask=attention_mask)
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
def test_generate_half_prec(self, test_name, model_id, config_cls, config_kwargs):
self._test_generate_half_prec(model_id, config_cls, config_kwargs)