mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-09-01 17:29:54 +00:00
[fix] sampler-ddp
This commit is contained in:
@@ -213,4 +213,5 @@ if __name__ == "__main__":
|
||||
|
||||
iter_per_epoch = len(train_loader)
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -262,4 +262,5 @@ if __name__ == "__main__":
|
||||
|
||||
iter_per_epoch = len(train_loader)
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -245,4 +245,5 @@ if __name__ == "__main__":
|
||||
|
||||
iter_per_epoch = len(train_loader)
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -199,4 +199,5 @@ if __name__ == "__main__":
|
||||
|
||||
iter_per_epoch = len(train_loader)
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -313,4 +313,5 @@ if __name__ == "__main__":
|
||||
model = DistributedDataParallel(model, device_ids=[ddp_local_rank])
|
||||
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
grpo_train_epoch(epoch, wandb)
|
||||
|
||||
@@ -205,4 +205,5 @@ if __name__ == "__main__":
|
||||
iter_per_epoch = len(train_loader)
|
||||
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -367,6 +367,7 @@ if __name__ == "__main__":
|
||||
old_actor_model.to(args.device)
|
||||
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
ppo_train_epoch(epoch, wandb, old_actor_model, ref_model, actor_scheduler, critic_scheduler)
|
||||
|
||||
if ddp:
|
||||
|
||||
@@ -197,4 +197,5 @@ if __name__ == "__main__":
|
||||
|
||||
iter_per_epoch = len(train_loader)
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
train_epoch(epoch, wandb)
|
||||
|
||||
@@ -364,4 +364,5 @@ if __name__ == "__main__":
|
||||
value_tracker = AutoAdaptiveValueTracker(rho_mode='kl', rho_const=0.9, D_half=0.06, clip_lower=0.5, clip_upper=0.96)
|
||||
|
||||
for epoch in range(args.epochs):
|
||||
train_sampler and train_sampler.set_epoch(epoch)
|
||||
spo_train_epoch(epoch, wandb, value_tracker)
|
||||
|
||||
Reference in New Issue
Block a user