ROCm: guard Trainer init patch against missing generated function

This commit is contained in:
Daniel Han-Chen 2026-02-11 05:04:48 +00:00
commit 7660d647e4

View file

@ -1930,8 +1930,23 @@ def patch_gradient_accumulation_fix(Trainer):
'if hasattr(unwrapped_model, "accepts_loss_kwargs") and False:',
init_function,
)
exec(init_function, globals())
Trainer.__init__ = _unsloth___init__
local_scope = {}
try:
exec(init_function, globals(), local_scope)
except Exception as _patch_error:
print(
"Unsloth: gradient accumulation init patch skipped due to "
f"source patch error: {_patch_error}"
)
local_scope = {}
_patched_init = local_scope.get("_unsloth___init__")
if _patched_init is None:
print(
"Unsloth: gradient accumulation init patch skipped because "
"_unsloth___init__ was not generated."
)
else:
Trainer.__init__ = _patched_init
def patch_tokenizer(model, tokenizer):
@ -2410,7 +2425,14 @@ def _prepare_model_for_qat(
except ImportError:
raise ImportError(TORCHAO_MSG)
group_size = 128
base_config = Int4WeightOnlyConfig(group_size = group_size)
try:
base_config = Int4WeightOnlyConfig(
group_size = group_size,
version = 2,
)
except TypeError:
# Older TorchAO versions do not support the version argument.
base_config = Int4WeightOnlyConfig(group_size = group_size)
filter_fn = (
lambda m, _: isinstance(m, torch.nn.Linear)
and m.in_features >= group_size
@ -2665,17 +2687,32 @@ def make_fast_generate_wrapper(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 DEVICE_TYPE == "hip":
# Allow GRPO notebooks to run on AMD without vLLM
print(
"Unsloth: `sampling_params` ignored because fast inference is "
"disabled on AMD."
)
kwargs.pop("sampling_params", None)
else:
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."
)
if DEVICE_TYPE == "hip":
print(
"Unsloth: `lora_request` ignored because fast inference is "
"disabled on AMD."
)
kwargs.pop("lora_request", None)
else:
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:
@ -2689,21 +2726,38 @@ def make_fast_generate_wrapper(original_generate):
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"
" )"
)
if DEVICE_TYPE == "hip":
model = getattr(original_generate, "__self__", None)
tokenizer = getattr(model, "_saved_temp_tokenizer", None)
if tokenizer is None:
raise ValueError(
"Unsloth: Passing text strings to `fast_generate` on AMD "
"requires a tokenizer attached to the model."
)
texts = [first_arg] if isinstance(first_arg, str) else list(first_arg)
tokens = tokenizer(
texts,
return_tensors="pt",
padding=True,
)
tokens = tokens.to(model.device)
return original_generate(**tokens, **kwargs)
else:
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)