From f659b55761b754d306bd140573493a6543cafd7f Mon Sep 17 00:00:00 2001 From: jingyaogong Date: Wed, 23 Sep 2026 00:07:30 +0800 Subject: [PATCH] [fix] residual step order --- trainer/train_grpo.py | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/trainer/train_grpo.py b/trainer/train_grpo.py index 21f5d9e..8a37dd6 100755 --- a/trainer/train_grpo.py +++ b/trainer/train_grpo.py @@ -144,7 +144,7 @@ def grpo_train_epoch(epoch, loader, iters, rollout_engine, ref_model, reward_mod loss = (policy_loss + aux_loss) / args.accumulation_steps # scalar loss.backward() - if step % args.accumulation_steps == 0: + if step % args.accumulation_steps == 0 or step == iters: if args.grad_clip > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() @@ -195,13 +195,6 @@ def grpo_train_epoch(epoch, loader, iters, rollout_engine, ref_model, reward_mod del prompt_inputs, outputs, completion_ids, per_token_logps, ref_per_token_logps del completions, rewards, grouped_rewards, mean_r, std_r, advantages, completion_mask, completion_pad_mask, prompt_lens, logp_pos - if step > start_step and step % args.accumulation_steps != 0: - if args.grad_clip > 0: - torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) - optimizer.step() - scheduler.step() - optimizer.zero_grad() - if __name__ == "__main__": parser = argparse.ArgumentParser(description="MiniMind GRPO (Group Relative Policy Optimization)")