From 721eee6a80e00313e79a2e09f6d9894628d7d41a Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 4 Sep 2025 03:25:40 -0700 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 95e9b79194..0f1fa2dbf6 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -532,7 +532,9 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # https://verl.readthedocs.io/en/latest/examples/config.html if trainer_file == "grpo_trainer": replacements = { - "beta" : 0.001, + "loss_type" : "bnpo", # Default GRPO paper + "beta" : 0.001, # Recommended as seen in verl + "auto_find_batch_size" : False, # Cannot work on GRPO } for k, v in replacements.items(): x = f"{k}( = [^,\n]{{1,}})?,\n" @@ -545,9 +547,9 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # Warn on too large or too small learning rate if " learning_rate" in call_args: learning_rate_check = \ - "if learning_rate < 1e-7: raise FloatingPointError(f'Unsloth: Your learning rate of `{learning_rate}` is too small and less than 1e-7! "\ + "if learning_rate < 1e-7: print(f'Unsloth: Your learning rate of `{learning_rate}` is too small and less than 1e-7! "\ "Consider increasing it, otherwise gradient updates will be close to 0!')\n"\ - "if learning_rate > 1: raise OverflowError(f'Unsloth: Your learning rate of `{learning_rate}` is way too larger > 1! "\ + "if learning_rate > 1: print(f'Unsloth: Your learning rate of `{learning_rate}` is way too larger > 1! "\ "Consider decreasing it to 1e-1, otherwise gradient updates will explode!')\n" extra_args += learning_rate_check pass @@ -614,9 +616,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): " print('Unsloth: The Dr GRPO paper recommends setting `scale_rewards` to False! Will override. Set it to `None` to force False.')\n"\ " scale_rewards = False\n"\ "elif loss_type.lower() == 'dapo':\n"\ - " print('Unsloth: The DAPO paper recommends `mask_truncated_completions = True`')\n"\ - " print('Unsloth: The DAPO paper recommends `epsilon_high = 0.28`')\n"\ - " print('Unsloth: The DAPO paper recommends setting `beta = 0.0` to remove the KL term')\n"\ + " if mask_truncated_completions != True:\n"\ + " print('Unsloth: The DAPO paper recommends `mask_truncated_completions = True`')\n"\ + " if epsilon_high != 0.28:\n"\ + " print('Unsloth: The DAPO paper recommends `epsilon_high = 0.28`')\n"\ + " if beta != 0.0:\n"\ + " print('Unsloth: The DAPO paper recommends setting `beta = 0.0` to remove the KL term')\n"\ " mask_truncated_completions = True\n"\ " epsilon_high = 0.28\n"\ " beta = 0.0\n"\