From 0c91dea4a28e1df5880499cedca361270a3ee025 Mon Sep 17 00:00:00 2001 From: "TOPSAIL\\78190" <781908858@qq.com> Date: Mon, 24 Aug 2026 12:43:50 +0800 Subject: [PATCH] [refactor] remove redundant rollout update_policy after torch.compile --- trainer/train_agent.py | 1 - trainer/train_grpo.py | 1 - trainer/train_ppo.py | 1 - 3 files changed, 3 deletions(-) diff --git a/trainer/train_agent.py b/trainer/train_agent.py index 445101c..e18aada 100644 --- a/trainer/train_agent.py +++ b/trainer/train_agent.py @@ -470,7 +470,6 @@ if __name__ == "__main__": if args.use_compile == 1: model = torch.compile(model) Logger('torch.compile enabled') - rollout_engine.update_policy(model) if dist.is_initialized(): model = DistributedDataParallel(model, device_ids=[local_rank]) rollout_engine.update_policy(model) diff --git a/trainer/train_grpo.py b/trainer/train_grpo.py index 200adb6..664e7e5 100755 --- a/trainer/train_grpo.py +++ b/trainer/train_grpo.py @@ -310,7 +310,6 @@ if __name__ == "__main__": if args.use_compile == 1: model = torch.compile(model) Logger('torch.compile enabled') - rollout_engine.update_policy(model) if dist.is_initialized(): model = DistributedDataParallel(model, device_ids=[local_rank]) rollout_engine.update_policy(model) diff --git a/trainer/train_ppo.py b/trainer/train_ppo.py index 0cf844d..f0427b6 100644 --- a/trainer/train_ppo.py +++ b/trainer/train_ppo.py @@ -426,7 +426,6 @@ if __name__ == "__main__": if args.use_compile == 1: actor_model = torch.compile(actor_model) Logger('torch.compile enabled') - rollout_engine.update_policy(actor_model) if dist.is_initialized(): actor_model = DistributedDataParallel(actor_model, device_ids=[local_rank]) critic_model = DistributedDataParallel(critic_model, device_ids=[local_rank])