Throw error when inferencing longer than max_popsition_embeddings (#1236)
* Throw error when inferencing longer than max_popsition_embeddings without rope scaling * Update llama.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com>
This commit is contained in:
parent
89d851b8bc
commit
9db2e12252
1 changed files with 9 additions and 0 deletions
|
|
@ -1376,6 +1376,15 @@ def _wrap_fast_inference(generate, device_type, dtype, model):
|
|||
@torch.inference_mode
|
||||
def _fast_generate(*args, **kwargs):
|
||||
|
||||
if hasattr(model, "config") and hasattr(model.config, "max_position_embeddings"):
|
||||
if "input_ids" in kwargs and kwargs["input_ids"] is not None and "max_new_tokens" in kwargs:
|
||||
if kwargs["input_ids"].shape[-1] + kwargs["max_new_tokens"] > model.config.max_position_embeddings:
|
||||
raise ValueError(
|
||||
f'Unsloth: input length {kwargs["input_ids"].shape[-1]} + max_new_tokens {kwargs["max_new_tokens"]} exceeds the maximum sequence length of {model.config.max_position_embeddings}!\n'\
|
||||
'You will need to do long context extension by increasing the `max_seq_length` in `FastLanguageModel.from_pretrained`.'
|
||||
)
|
||||
pass
|
||||
|
||||
# Set a flag for generation!
|
||||
internal_model = model
|
||||
while hasattr(internal_model, "model"):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue