mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
fix generate because of recent transformers release
This commit is contained in:
+8
-17
@@ -539,16 +539,16 @@ class PeftModelForCausalLM(PeftModel):
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
batch_size = kwargs["input_ids"].shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
kwargs["past_key_values"] = past_key_values
|
||||
return self.base_model.generate(**kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
model_kwargs["past_key_values"] = kwargs.get("past", None) or kwargs.get("past_key_values", None)
|
||||
if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
batch_size = model_kwargs["input_ids"].shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
model_kwargs["past_key_values"] = past_key_values
|
||||
return model_kwargs
|
||||
|
||||
|
||||
@@ -682,25 +682,16 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
batch_size = kwargs["input_ids"].shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
kwargs["past_key_values"] = past_key_values
|
||||
return self.base_model.generate(**kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
model_kwargs["past_key_values"] = kwargs.get("past", None) or kwargs.get("past_key_values", None)
|
||||
return model_kwargs
|
||||
|
||||
def _prepare_encoder_decoder_kwargs_for_generation(self, inputs_tensor, model_kwargs, model_input_name=None):
|
||||
past_key_values = model_kwargs.get("past_key_values", None)
|
||||
model_kwargs["past_key_values"] = None
|
||||
model_kwargs = self.base_model_prepare_encoder_decoder_kwargs_for_generation(
|
||||
inputs_tensor, model_kwargs, model_input_name
|
||||
)
|
||||
model_kwargs["past_key_values"] = past_key_values
|
||||
if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
batch_size = model_kwargs["decoder_input_ids"].shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
model_kwargs["past_key_values"] = past_key_values
|
||||
return model_kwargs
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user