mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-09-27 05:17:22 +00:00
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:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user