Update rl.py
This commit is contained in:
parent
8a0d710f22
commit
ecb22588d1
1 changed files with 22 additions and 0 deletions
|
|
@ -481,6 +481,28 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
extra_args += num_proc_check
|
||||
pass
|
||||
|
||||
# Check for loss_type = dr_grpo and scale_rewards for GRPO
|
||||
if "loss_type" in call_args and "scale_rewards" in call_args:
|
||||
check_dr_grpo = \
|
||||
"if loss_type.lower() == 'dr_grpo':\n"\
|
||||
" loss_type = 'dr_grpo'\n"\
|
||||
"elif loss_type.lower() == 'dapo':\n"\
|
||||
" loss_type = 'dapo'\n"\
|
||||
"if loss_type.lower() == 'dr_grpo':\n"\
|
||||
" if scale_rewards == None:\n"\
|
||||
" scale_rewards = True\n"\
|
||||
" elif scale_rewards == True:\n"\
|
||||
" print('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('The DAPO paper recommends `mask_truncated_completions = True`')\n"\
|
||||
" print('The DAPO paper recommends `epsilon_high = 0.28`')\n"\
|
||||
" mask_truncated_completions = True\n"\
|
||||
" epsilon_high = 0.28\n"\
|
||||
"\n"
|
||||
extra_args += check_dr_grpo
|
||||
pass
|
||||
|
||||
# Edit config with anything extra
|
||||
if trainer_file in RL_CONFIG_CHANGES:
|
||||
process_extra_args = RL_CONFIG_CHANGES[trainer_file]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue