Update rl_replacements.py
This commit is contained in:
parent
caca33f401
commit
02b45397d6
1 changed files with 6 additions and 2 deletions
|
|
@ -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"\
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue