ITADN

Does PPOTrainer use mini_batch_size to update parameters?

#5645Closedzhiang28 创建于 2026-04-26
Z
zhiang28commented
I was reviewing the code for PPOTrainer in trl/experimental/ppo/ppo_trainer.py. Specifically, between lines 804 and 882, it appears that optimizer.step() is called based on local_batch_size (per_device_train_batch_size * gradient_accumulation_steps) rather than mini_batch_size. ``` for ppo_epoch_idx in range(args.num_ppo_epochs): b_inds = np.random.permutation(args.local_batch_size) minibatch_idx = 0 for mini_batch_start in range(0, args.local_batch_size, args.local_mini_batch_size): mini_batch_end = mini_batch_start + args.local_mini_batch_size mini_batch_inds = b_inds[mini_batch_start:mini_batch_end] gradient_accumulation_idx = 0 for micro_batch_start in range(0, args.local_mini_batch_size, args.per_device_train_batch_size): with accelerator.accumulate(model): micro_batch_end = micro_batch_start + args.per_device_train_batch_size micro_batch_inds = mini_batch_inds[micro_batch_start:micro_batch_end] ....... accelerator.backward(loss) optimizer.step() optimizer.zero_grad() with torch.no_grad(): ......... gradient_accumulation_idx += 1 minibatch_idx += 1 # del everything and empty cache # fmt: off del ( output, vpred_temp, logits, new_logprobs, vpred, vpredclipped, vf_losses1, vf_losses2, vf_loss, vf_clipfrac, logprobs_diff, ratio, pg_losses, pg_losses2, pg_loss_max, pg_loss, loss, pg_clipfrac, prob_dist, entropy, approxkl, mb_return, mb_advantage, mb_values, mb_responses, mb_query_responses, mb_logprobs, ) # fmt: on empty_cache()
关闭于 2026-04-29 2 条评论