From 02b45397d60f00d22ab9ea55b16b257e0452249f Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Fri, 14 Feb 2025 16:01:29 -0800 Subject: [PATCH] Update rl_replacements.py --- unsloth/models/rl_replacements.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index c7fdb4cbde..682a35ed1c 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -248,8 +248,12 @@ RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer_compute_loss) # https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py#L356 # TRL warns if batch size is not a multiple of num_generations -> fix this. def grpo_trainer_fix_batch_size(RLTrainer_source, RLConfig_source): - if "divisible by the number of generations" not in RLTrainer_source: return "" - if "num_generations" not in RLConfig_source: return "" + if "divisible by the number of generations" not in RLTrainer_source: + print(RLTrainer_source) + return "" + if "num_generations" not in RLConfig_source: + print(RLConfig_source) + return "" check_batch_size = \ "div = per_device_train_batch_size // num_generations\n"\