From 3b5c0d3c74a353a2af2bd721fab154176ed45964 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 17 Feb 2025 23:57:57 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 231dbe7765..7a90b81157 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -442,6 +442,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # Selective log softmax selective_log_softmax_code = inspect.getsource(selective_log_softmax) + # Trainer kwargs + comma = "" if RLTrainer_call_args.endswith(",") else "," + unsloth_extra_args = comma + \ + "vllm_sampling_params = vllm_sampling_params,\n"\ + "unsloth_num_chunks = unsloth_num_chunks, **kwargs" + # Get final source code RLTrainer_source = RLTrainer_replacement.format( RLTrainer_name = RLTrainer_name, @@ -449,7 +455,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): RLTrainer_arguments = RLTrainer_arguments, RLTrainer_extra_args = RLTrainer_extra_args, RLTrainer_call_args = RLTrainer_call_args, - RLTrainer_kwargs = ",**kwargs"[1 if RLTrainer_call_args.endswith(",") else 0:], + RLTrainer_kwargs = unsloth_extra_args, RLConfig_name = RLConfig_name, __RLConfig_doc__ = __RLConfig_doc__,