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