Update rl_replacements.py

This commit is contained in:
Daniel Han 2025-02-14 16:03:13 -08:00
commit 42a576a220

View file

@ -248,18 +248,14 @@ 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:
print(RLTrainer_source)
return ""
if "num_generations" not in RLConfig_source:
print(RLConfig_source)
return ""
if "divisible by the number of generations" not in RLTrainer_source: return ""
if "num_generations" not in RLConfig_source: return ""
check_batch_size = \
"div = per_device_train_batch_size // num_generations\n"\
"if div * num_generations != per_device_train_batch_size:\n"\
" print('Unsloth: We know expect `per_device_train_batch_size` to be a multiple of `num_generations`.\\n'\\"\
" 'We will change the batch size of ' + per_device_train_batch_size + ' to the `num_generations` of ' + num_generations')\n"\
" print('Unsloth: We know expect `per_device_train_batch_size` to be a multiple of `num_generations`.\\n"\
"We will change the batch size of ' + str(per_device_train_batch_size) + ' to the `num_generations` of ' + str(num_generations)')\n"\
" per_device_train_batch_size = num_generations\n"
return check_batch_size
pass