mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
changed: 1. replace base_model.prepare_inputs_for_generation and base_model._prepare_encoder_decoder_kwargs_for_generation temporarily
This commit is contained in:
+67
-43
@@ -513,7 +513,6 @@ class PeftModelForCausalLM(PeftModel):
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -576,28 +575,38 @@ class PeftModelForCausalLM(PeftModel):
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
def generate(self, **kwargs):
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
return self.base_model.generate(**kwargs)
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
try:
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
if "input_ids" not in kwargs:
|
||||
raise ValueError("input_ids must be provided for Peft model generation")
|
||||
if kwargs.get("attention_mask", None) is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(
|
||||
kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens
|
||||
).to(kwargs["input_ids"].device)
|
||||
kwargs["attention_mask"] = torch.cat((prefix_attention_mask, kwargs["attention_mask"]), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Position ids are not supported for parameter efficient tuning. Ignoring position ids."
|
||||
)
|
||||
kwargs["position_ids"] = None
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Token type ids are not supported for parameter efficient tuning. Ignoring token type ids"
|
||||
)
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
except:
|
||||
self.base_model.prepare_inputs_for_generation = self.base_model_prepare_inputs_for_generation
|
||||
raise
|
||||
else:
|
||||
if "input_ids" not in kwargs:
|
||||
raise ValueError("input_ids must be provided for Peft model generation")
|
||||
if kwargs.get("attention_mask", None) is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(
|
||||
kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens
|
||||
).to(kwargs["input_ids"].device)
|
||||
kwargs["attention_mask"] = torch.cat((prefix_attention_mask, kwargs["attention_mask"]), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.")
|
||||
kwargs["position_ids"] = None
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Token type ids are not supported for parameter efficient tuning. Ignoring token type ids"
|
||||
)
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
return self.base_model.generate(**kwargs)
|
||||
self.base_model.prepare_inputs_for_generation = self.base_model_prepare_inputs_for_generation
|
||||
return outputs
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
@@ -641,13 +650,9 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
self.base_model_prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self._prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -740,24 +745,43 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
)
|
||||
|
||||
def generate(self, **kwargs):
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
return self.base_model.generate(**kwargs)
|
||||
else:
|
||||
if "input_ids" not in kwargs:
|
||||
raise ValueError("input_ids must be provided for Peft model generation")
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.")
|
||||
kwargs["position_ids"] = None
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Token type ids are not supported for parameter efficient tuning. Ignoring token type ids"
|
||||
)
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
return self.base_model.generate(**kwargs)
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self._prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
try:
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if "input_ids" not in kwargs:
|
||||
raise ValueError("input_ids must be provided for Peft model generation")
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Position ids are not supported for parameter efficient tuning. Ignoring position ids."
|
||||
)
|
||||
kwargs["position_ids"] = None
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
warnings.warn(
|
||||
"Token type ids are not supported for parameter efficient tuning. Ignoring token type ids"
|
||||
)
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
except:
|
||||
self.base_model.prepare_inputs_for_generation = self.base_model_prepare_inputs_for_generation
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self.base_model_prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
raise
|
||||
else:
|
||||
self.base_model.prepare_inputs_for_generation = self.base_model_prepare_inputs_for_generation
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self.base_model_prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
return outputs
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
|
||||
Reference in New Issue
Block a user