From 8266e2ee4fb552aa530e50f98dd5412bc121ebc5 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 10:01:47 +0000 Subject: [PATCH 1/2] fix half precision forward --- src/peft/peft_model.py | 17 +++++++++++++++++ src/peft/tuners/lora.py | 35 ++++++++++++++++++++++++----------- tests/test_peft_model.py | 23 +++++++++++++++++++++++ 3 files changed, 64 insertions(+), 11 deletions(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index f9573bb..a39e046 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -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: diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 51cd56f..a1fa310 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -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(): diff --git a/tests/test_peft_model.py b/tests/test_peft_model.py index 4280ff3..53c0ab9 100644 --- a/tests/test_peft_model.py +++ b/tests/test_peft_model.py @@ -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) From f35b20a845f38af258682d71e7e30253aeb59b2b Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Fri, 7 Apr 2023 10:48:22 +0000 Subject: [PATCH 2/2] add and fix tests --- src/peft/peft_model.py | 16 +++++++++++++++- src/peft/tuners/lora.py | 8 ++++++++ tests/test_decoder_models.py | 4 ++++ tests/test_encoder_decoder_models.py | 4 ++++ tests/testing_common.py | 22 ++++++++++++++++++++++ 5 files changed, 53 insertions(+), 1 deletion(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index c4881b4..652f4aa 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -85,6 +85,7 @@ class PeftModel(PushToHubMixin, torch.nn.Module): self.peft_config = {} self.active_adapter = adapter_name self.peft_type = peft_config.peft_type + self.base_model_torch_dtype = getattr(model, "dtype", None) if not isinstance(peft_config, PromptLearningConfig): self.peft_config[adapter_name] = peft_config self.base_model = PEFT_TYPE_TO_MODEL_MAPPING[peft_config.peft_type]( @@ -93,7 +94,6 @@ class PeftModel(PushToHubMixin, torch.nn.Module): else: self.add_adapter(adapter_name, peft_config) - def save_pretrained(self, save_directory, **kwargs): r""" This function saves the adapter model and the adapter configuration files to a directory, so that it can be @@ -967,7 +967,21 @@ class PeftModelForSeq2SeqLM(PeftModel): if model_kwargs["past_key_values"] is None and peft_config.peft_type == PeftType.PREFIX_TUNING: batch_size = model_kwargs["decoder_input_ids"].shape[0] past_key_values = self.get_prompt(batch_size) + 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 + return model_kwargs diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 90e4a23..3b70dbb 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -459,6 +459,8 @@ class Linear(nn.Linear, LoraLayer): self.merged = False def forward(self, x: torch.Tensor): + previous_dtype = x.dtype + if self.active_adapter not in self.lora_A.keys(): return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) if self.disable_adapters: @@ -467,6 +469,9 @@ class Linear(nn.Linear, LoraLayer): result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) elif self.r[self.active_adapter] > 0 and not self.merged: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + + x = x.to(self.lora_A[self.active_adapter].weight.dtype) + result += ( self.lora_B[self.active_adapter]( self.lora_A[self.active_adapter](self.lora_dropout[self.active_adapter](x)) @@ -475,6 +480,9 @@ class Linear(nn.Linear, LoraLayer): ) else: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + + result = result.to(previous_dtype) + return result diff --git a/tests/test_decoder_models.py b/tests/test_decoder_models.py index 209b4df..cdbf56b 100644 --- a/tests/test_decoder_models.py +++ b/tests/test_decoder_models.py @@ -83,3 +83,7 @@ class PeftDecoderModelTester(unittest.TestCase, PeftCommonTester): @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) + + @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) diff --git a/tests/test_encoder_decoder_models.py b/tests/test_encoder_decoder_models.py index cdf9571..974e214 100644 --- a/tests/test_encoder_decoder_models.py +++ b/tests/test_encoder_decoder_models.py @@ -86,3 +86,7 @@ class PeftEncoderDecoderModelTester(unittest.TestCase, PeftCommonTester): @parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID, filter_params_func=skip_non_lora_or_pt)) def test_generate(self, test_name, model_id, config_cls, config_kwargs): self._test_generate(model_id, config_cls, config_kwargs) + + @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) diff --git a/tests/testing_common.py b/tests/testing_common.py index cfe6cf2..0d0d169 100644 --- a/tests/testing_common.py +++ b/tests/testing_common.py @@ -297,3 +297,25 @@ class PeftCommonTester: with self.assertRaises(TypeError): # check if `generate` raises an error if no positional arguments are passed _ = model.generate(inputs["input_ids"]) + + def _test_generate_half_prec(self, model_id, config_cls, config_kwargs): + if config_cls not in (LoraConfig, PrefixTuningConfig): + return + + model = self.transformers_class.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)