Update rl_replacements.py
This commit is contained in:
parent
7e4f7f92a5
commit
42a576a220
1 changed files with 4 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue