ROCm: guard Trainer init patch against missing generated function
This commit is contained in:
parent
d02190bd12
commit
7660d647e4
1 changed files with 81 additions and 27 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue