Merge pull request #821 from lycDragonSlayer/refactor/remove-redundant-update-policy

[refactor] remove redundant rollout update_policy after torch.compile
This commit is contained in:
jingyaogong
2026-09-20 20:13:29 +08:00
committed by GitHub
3 changed files with 0 additions and 3 deletions
-1
View File
@@ -471,7 +471,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():
# 同 train_ppo:RoPE buffer 各 rank 一致,每步广播纯属浪费
model = DistributedDataParallel(model, device_ids=[local_rank], broadcast_buffers=False)
-1
View File
@@ -312,7 +312,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():
# 同 train_ppo:RoPE buffer 各 rank 一致,每步广播纯属浪费
model = DistributedDataParallel(model, device_ids=[local_rank], broadcast_buffers=False)
-1
View File
@@ -427,7 +427,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():
# freqs_cos/freqs_sin 各 rank 由 config 确定性算出,默认每步广播一次纯属浪费
actor_model = DistributedDataParallel(actor_model, device_ids=[local_rank], broadcast_buffers=False)