Add helpful error messages for fast_generate when fast_inference=False

When users load a model with fast_inference=False but then try to use
vLLM-style arguments with fast_generate, they previously got confusing
errors. This adds a wrapper that detects common mistakes and provides
helpful guidance:

- Using sampling_params: explains to use HF generate args instead
- Using lora_request: explains LoRA weights are already merged
- Passing text strings: shows how to tokenize input first

Changes:
- Add make_fast_generate_wrapper to _utils.py
- Apply wrapper in llama.py when fast_inference=False
- Apply wrapper in vision.py when fast_inference=False
This commit is contained in:
danielhanchen 2026-01-02 13:58:08 +00:00
commit 91d911ff81
3 changed files with 58 additions and 2 deletions

View file

@ -73,6 +73,7 @@ __all__ = [
"verify_fp8_support_if_applicable",
"_get_inference_mode_context_manager",
"hf_login",
"make_fast_generate_wrapper",
]
import torch
@ -2378,3 +2379,58 @@ def hf_login(token: Optional[str] = None) -> Optional[str]:
except Exception as e:
logger.info(f"Failed to login to huggingface using token with error: {e}")
return token
def make_fast_generate_wrapper(original_generate):
"""
Creates a wrapper around model.generate that checks for incorrect
vLLM-style usage when fast_inference=False.
"""
@functools.wraps(original_generate)
def _fast_generate_wrapper(*args, **kwargs):
# Check for vLLM-specific arguments
if "sampling_params" in kwargs:
raise ValueError(
"Unsloth: `sampling_params` is only supported when `fast_inference=True` (vLLM). "
"Since `fast_inference=False`, use HuggingFace generate arguments instead:\n"
" model.fast_generate(**tokens.to('cuda'), max_new_tokens=64, temperature=1.0, top_p=0.95)"
)
if "lora_request" in kwargs:
raise ValueError(
"Unsloth: `lora_request` is only supported when `fast_inference=True` (vLLM). "
"Since `fast_inference=False`, LoRA weights are already merged into the model."
)
# Check if first positional argument is a string or list of strings
if len(args) > 0:
first_arg = args[0]
is_string_input = False
if isinstance(first_arg, str):
is_string_input = True
elif isinstance(first_arg, (list, tuple)) and len(first_arg) > 0:
if isinstance(first_arg[0], str):
is_string_input = True
if is_string_input:
raise ValueError(
"Unsloth: Passing text strings to `fast_generate` is only supported "
"when `fast_inference=True` (vLLM). Since `fast_inference=False`, you must "
"tokenize the input first:\n\n"
" messages = tokenizer.apply_chat_template(\n"
" [{\"role\": \"user\", \"content\": \"Your prompt here\"}],\n"
" tokenize=True, add_generation_prompt=True,\n"
" return_tensors=\"pt\", return_dict=True\n"
" )\n"
" output = model.fast_generate(\n"
" **messages.to('cuda'),\n"
" max_new_tokens=64,\n"
" temperature=1.0,\n"
" )"
)
# Call original generate
return original_generate(*args, **kwargs)
return _fast_generate_wrapper

View file

@ -2326,7 +2326,7 @@ class FastLlamaModel:
attn_implementation = "eager",
**kwargs,
)
model.fast_generate = model.generate
model.fast_generate = make_fast_generate_wrapper(model.generate)
model.fast_generate_batches = None
else:
from unsloth_zoo.vllm_utils import (

View file

@ -673,7 +673,7 @@ class FastBaseModel:
**kwargs,
)
if hasattr(model, "generate"):
model.fast_generate = model.generate
model.fast_generate = make_fast_generate_wrapper(model.generate)
model.fast_generate_batches = error_out_no_vllm
if offload_embedding:
if bool(