Repository navigation
Expand file tree
/
Copy pathtrain.py
More file actions
133 lines (108 loc) · 5.41 KB
/
Copy pathtrain.py
File metadata and controls
133 lines (108 loc) · 5.41 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
import logging
import ray
from slime.data.checkpoint import save_checkpoint
from slime.observability.logging_utils import configure_logger, finish_tracking, init_tracking
from slime.ray.placement_group import create_rollout_manager, create_training_models
from slime.utils.arguments import parse_args
from slime.utils.cleanup import Cleanup
from slime.utils.misc import should_run_periodic_action
def train(args, pgs, rollout_manager, num_rollout_per_epoch, restore_plan):
release_train = args.release_train
actor_model, critic_model = create_training_models(args, pgs, rollout_manager)
if args.offload_rollout and not release_train:
ray.get(rollout_manager.onload_weights.remote())
# Always push actor weights to rollout once weights are loaded.
actor_model.update_weights()
# Reattached async producers stay paused until restored weights are installed.
# This startup notification resumes them once, before entering the train loop.
ray.get(rollout_manager.training_ready.remote())
if args.check_weight_update_equal:
ray.get(rollout_manager.check_weights.remote(action="compare"))
if args.offload_rollout:
ray.get(rollout_manager.onload_kv.remote())
# special case for eval-only
if args.num_rollout == 0 and args.eval_interval is not None:
ray.get(rollout_manager.eval.remote(rollout_id=0))
def offload_train(actor_trains_this_step):
# Each model auto-offloads after train() when offload_train is set,
# so we only need clear_memory for the non-offload case.
if not args.offload_train:
if not args.use_critic or actor_trains_this_step:
actor_model.clear_memory()
else:
critic_model.clear_memory()
# train loop.
for rollout_id in range(args.start_rollout_id, args.num_rollout):
if args.eval_interval is not None and rollout_id == 0 and not args.skip_eval_before_train:
ray.get(rollout_manager.eval.remote(rollout_id))
rollout_data_ref = ray.get(rollout_manager.generate.remote(rollout_id))
if args.offload_rollout:
ray.get(rollout_manager.offload.remote())
if release_train:
actor_model.create()
actor_trains = (not args.use_critic) or rollout_id >= args.num_critic_only_steps
if args.use_critic:
value_refs = critic_model.async_train(rollout_id, rollout_data_ref)
if actor_trains:
ray.get(actor_model.async_train(rollout_id, rollout_data_ref, external_data=value_refs))
else:
ray.get(value_refs)
else:
ray.get(actor_model.async_train(rollout_id, rollout_data_ref))
# Runtime completion releases queue capacity, without claiming that the
# model/optimizer or this data progress have a durable joint checkpoint.
ray.get(rollout_manager.training_completed.remote(rollout_id))
if release_train or should_run_periodic_action(
rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout
):
save_checkpoint(
args,
rollout_id,
actor_model,
critic_model,
rollout_manager,
actor_trains=actor_trains,
restore_plan=restore_plan,
)
offload_train(actor_trains)
if args.offload_rollout and not release_train:
ray.get(rollout_manager.onload_weights.remote())
was_paused = ray.get(rollout_manager.pause_rollout_admission.remote())
actor_model.update_weights()
# Final evaluation uses the synchronized engines directly. Keep training
# producers paused when no later rollout will consume their new work.
if rollout_id + 1 < args.num_rollout:
ray.get(rollout_manager.resume_rollout_admission.remote(was_paused))
if args.offload_rollout:
ray.get(rollout_manager.onload_kv.remote())
if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch):
ray.get(rollout_manager.eval.remote(rollout_id))
def main(args, restore_plan):
"""Own the attempt's exit policy; actors only clean up their own resources."""
with Cleanup() as cleanup:
init_tracking(args)
startup = None
try:
# Create or reattach rollout resources, then run this training attempt.
startup = create_rollout_manager(args, restore_plan=restore_plan)
train(
startup.args, startup.placements, startup.manager, startup.num_rollout_per_epoch, startup.restore_plan
)
except BaseException:
# Failures and interruptions retain serving and replay data for a manual
# restart. Detach this driver's trainers instead of disposing the session.
if startup is not None:
try:
startup.close(failed=True)
except Exception:
logging.getLogger(__name__).exception("Failed to clean up training attempt")
raise
else:
startup.close(failed=False)
finally:
# Finish this driver's tracking on success, startup failure, or interruption.
cleanup.run("finish driver tracking", finish_tracking, args)
if __name__ == "__main__":
args, restore_plan = parse_args(return_restore_plan=True)
configure_logger()
main(args, restore_plan)