Fix GRPO (#2787)
* Update _utils.py * Update _utils.py * versioning * Update _utils.py * Update _utils.py * Update _utils.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update vision.py * HF Transfer * fix(utils): add missing importlib import to fix NameError (#2134) This commit fixes a NameError that occurs when `importlib` is referenced in _utils.py without being imported, especially when UNSLOTH_USE_MODELSCOPE=1 is enabled. By adding the missing import statement, the code will no longer throw a NameError. * Add QLoRA Train and Merge16bit Test (#2130) * add reference and unsloth lora merging tests * add test / dataset printing to test scripts * allow running tests from repo root * add qlora test readme * more readme edits * ruff formatting * additional readme comments * forgot to add actual tests * add apache license * Update pyproject.toml * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Revert * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Bug fix * Update mapper.py * check SDPA for Mistral 3, Pixtral * Update vision.py * Versioning * Update rl_replacements.py * Update README.md * add model registry * move hf hub utils to unsloth/utils * refactor global model info dicts to dataclasses * fix dataclass init * fix llama registration * remove deprecated key function * start registry reog * add llama vision * quant types -> Enum * remap literal quant types to QuantType Enum * add llama model registration * fix quant tag mapping * add qwen2.5 models to registry * add option to include original model in registry * handle quant types per model size * separate registration of base and instruct llama3.2 * add QwenQVQ to registry * add gemma3 to registry * add phi * add deepseek v3 * add deepseek r1 base * add deepseek r1 zero * add deepseek distill llama * add deepseek distill models * remove redundant code when constructing model names * add mistral small to registry * rename model registration methods * rename deepseek registration methods * refactor naming for mistral and phi * add global register models * refactor model registration tests for new registry apis * add model search method * remove deprecated registration api * add quant type test * add registry readme * make llama registration more specific * clear registry when executing individual model registration file * more registry readme updates * Update _auto_install.py * Llama4 * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Synthetic data * Update mapper.py * Xet and Synthetic * Update synthetic.py * Update loader.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * logits / temperature * Update rl_replacements.py * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py --------- Co-authored-by: naliazheli <nalia0316@gmail.com> Co-authored-by: jeromeku <jerome.ku@gmail.com> Co-authored-by: Jack Shi Wei Lun <87535974+jackswl@users.noreply.github.com> Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com>
This commit is contained in:
parent
c340fa3ac8
commit
c81b1864be
3 changed files with 37 additions and 4 deletions
|
|
@ -37,7 +37,7 @@ triton = [
|
|||
]
|
||||
|
||||
huggingface = [
|
||||
"unsloth_zoo>=2025.6.3",
|
||||
"unsloth_zoo>=2025.6.4",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2",
|
||||
|
|
@ -381,7 +381,7 @@ colab-ampere-torch220 = [
|
|||
"flash-attn>=2.6.3",
|
||||
]
|
||||
colab-new = [
|
||||
"unsloth_zoo>=2025.6.3",
|
||||
"unsloth_zoo>=2025.6.4",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2",
|
||||
|
|
|
|||
|
|
@ -486,6 +486,21 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
arguments = re.sub(x, y, arguments)
|
||||
pass
|
||||
|
||||
# Fix GRPO beta default as 0.001 TRL used to be 0.04, now 0.00!
|
||||
# https://github.com/huggingface/trl/pull/3516
|
||||
# https://verl.readthedocs.io/en/latest/examples/config.html
|
||||
if trainer_file == "grpo_trainer":
|
||||
replacements = {
|
||||
"beta" : 0.001,
|
||||
}
|
||||
for k, v in replacements.items():
|
||||
x = f"{k}( = [^,\n]{{1,}})?,\n"
|
||||
y = f"'{v}'" if type(v) is str else f"{v}"
|
||||
y = f"{k} = {y},\n"
|
||||
arguments = re.sub(x, y, arguments)
|
||||
pass
|
||||
pass
|
||||
|
||||
# Warn on too large or too small learning rate
|
||||
if " learning_rate" in call_args:
|
||||
learning_rate_check = \
|
||||
|
|
@ -553,6 +568,17 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
extra_args += check_num_generations
|
||||
pass
|
||||
|
||||
# Check temperature must not be <= 0. Also stop if >= 10
|
||||
if "temperature" in call_args:
|
||||
check_temperature = \
|
||||
"if temperature <= 0:\n"\
|
||||
" raise MathError('Unsloth: Please set a positive non-zero temperature since your results will be wrong.')\n"\
|
||||
"elif temperature >= 10:\n"\
|
||||
" raise MathError('Unsloth: Please set a positive non-zero temperature less than 10, since sampling will be quite erratic.')\n"\
|
||||
"\n"
|
||||
extra_args += check_temperature
|
||||
pass
|
||||
|
||||
# Edit config with anything extra
|
||||
if trainer_file in RL_CONFIG_CHANGES:
|
||||
process_extra_args = RL_CONFIG_CHANGES[trainer_file]
|
||||
|
|
|
|||
|
|
@ -269,6 +269,8 @@ def grpo_trainer__get_per_token_logps(function_name, function):
|
|||
# See https://github.com/huggingface/trl/issues/2770
|
||||
# logits = logits[:, -logits_to_keep:]
|
||||
# return logits
|
||||
# See https://huggingface.co/blog/the_n_implementation_details_of_rlhf_with_ppo#policy-training-implementation-details
|
||||
# logits = logits / self.temperature
|
||||
# logps = selective_log_softmax(logits, input_ids)
|
||||
|
||||
# row_indices, col_indices = torch.where(logps < -20)
|
||||
|
|
@ -325,7 +327,6 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
else:
|
||||
ref_per_token_logps = None
|
||||
# per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1
|
||||
|
||||
# x - x.detach() allows for preserving gradients from x
|
||||
advantages = inputs["advantages"]
|
||||
# per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)
|
||||
|
|
@ -335,10 +336,13 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
old_hidden_states = inputs["old_per_token_logps"]
|
||||
else:
|
||||
old_hidden_states = None
|
||||
|
||||
input_ids = input_ids[:, -logits_to_keep:]
|
||||
if per_token_logps is not None:
|
||||
|
||||
ref_per_token_logps = ref_per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
if ref_per_token_logps is not None:
|
||||
ref_per_token_logps = ref_per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
|
||||
per_token_logps = per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
|
||||
loss, completion_length, mean_kl = grpo_compute_loss_slow(
|
||||
|
|
@ -354,6 +358,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
epsilon_high = self.epsilon_high,
|
||||
max_completion_length = self.args.max_completion_length,
|
||||
delta = self.args.delta,
|
||||
temperature = self.args.temperature,
|
||||
)
|
||||
else:
|
||||
if hasattr(self.args, "loss_type"):
|
||||
|
|
@ -370,6 +375,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
epsilon_high = self.epsilon_high,
|
||||
max_completion_length = self.args.max_completion_length,
|
||||
delta = self.args.delta,
|
||||
temperature = self.args.temperature,
|
||||
)
|
||||
else:
|
||||
# to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17
|
||||
|
|
@ -381,6 +387,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
advantages,
|
||||
old_hidden_states,
|
||||
n_chunks = self.args.unsloth_num_chunks,
|
||||
temperature = self.args.temperature,
|
||||
)
|
||||
|
||||
# Log the metrics
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue