mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
fix half precision forward
This commit is contained in:
@@ -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
@@ -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():
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user