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:
Datta Nimmaturi 2024-11-07 01:52:08 +05:30 committed by GitHub
commit 9db2e12252

View file

@ -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"):