forked from yyht/miles-values
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
148 lines (124 loc) · 6.37 KB
/
Copy pathtrain.py
File metadata and controls
148 lines (124 loc) · 6.37 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
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
import asyncio
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH, GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_WEIGHTS
from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models
from miles.utils.arguments import parse_args
from miles.utils.async_utils import eager_create_task
from miles.utils.logging_utils import configure_logger
from miles.utils.misc import should_run_periodic_action
from miles.utils.tracking_utils import finish_tracking, init_tracking
async def train(args):
configure_logger()
# allocate the GPUs
pgs = create_placement_groups(args)
init_tracking(args)
# create the rollout manager, with sglang engines inside.
# need to initialize rollout manager first to calculate num_rollout
rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"])
# create the actor and critic models
actor_model, critic_model = await create_training_models(args, pgs, rollout_manager)
if args.offload_rollout:
await rollout_manager.onload_weights.remote()
# always update weight first so that sglang has the loaded weights from training.
await actor_model.update_weights()
if args.check_weight_update_equal:
await rollout_manager.check_weights.remote(action="compare")
if args.offload_rollout:
await rollout_manager.onload_kv.remote()
# special case for eval-only
if args.num_rollout == 0 and args.eval_interval is not None:
await rollout_manager.eval.remote(rollout_id=0)
async def offload_train():
if getattr(args, "colocate_critic", False):
# Colocate-critic: each model auto-sleeps inside train(),
# so no explicit offload needed here.
pass
elif args.offload_train:
if args.use_critic:
await critic_model.offload()
if rollout_id >= args.num_critic_only_steps:
await actor_model.offload()
else:
await actor_model.offload()
else:
await actor_model.clear_memory()
async def save(rollout_id):
if (not args.use_critic) or (rollout_id >= args.num_critic_only_steps):
await actor_model.save_model(
rollout_id,
force_sync=rollout_id == args.num_rollout - 1,
)
if args.use_critic:
await critic_model.save_model(
rollout_id,
force_sync=rollout_id == args.num_rollout - 1,
)
if args.rollout_global_dataset:
await rollout_manager.save.remote(rollout_id)
# train loop.
# note that for async training, one can change the position of the sync operation(ray.get).
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:
await rollout_manager.eval.remote(rollout_id)
rollout_data_ref = await rollout_manager.generate.remote(rollout_id)
if args.offload_rollout:
offload_tags = [GPU_MEMORY_TYPE_CUDA_GRAPH]
if "kv_cache" in args.offload_rollout_level:
offload_tags.append(GPU_MEMORY_TYPE_KV_CACHE)
if "weight" in args.offload_rollout_level:
offload_tags.append(GPU_MEMORY_TYPE_WEIGHTS)
await rollout_manager.offload.remote(tags=offload_tags)
if args.use_critic:
if getattr(args, "colocate_critic", False):
# Colocate-critic: submit critic first (returns Ray ObjectRefs),
# then submit actor with those refs as external_data.
# Ray auto-resolves the dependency: actor won't wake_up until
# critic has finished and slept, preventing GPU OOM.
value_refs = critic_model.submit_train(rollout_id, rollout_data_ref)
if rollout_id >= args.num_critic_only_steps:
actor_refs = actor_model.submit_train(
rollout_id, rollout_data_ref, external_data=value_refs,
)
await asyncio.gather(*actor_refs)
else:
await asyncio.gather(*value_refs)
else:
# Separated-critic: critic and actor train concurrently, sync via NCCL.
critic_task = await eager_create_task(critic_model.train(rollout_id, rollout_data_ref))
if rollout_id >= args.num_critic_only_steps:
await actor_model.train(rollout_id, rollout_data_ref)
await critic_task
else:
await actor_model.train(rollout_id, rollout_data_ref)
if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout):
await save(rollout_id)
await offload_train()
# Sync actor weights to rollout engines.
# During critic-only steps, actor weights are unchanged but we still
# need to onload → update_weights → onload_kv when colocate is enabled,
# to ensure CUDA graphs are rebuilt AFTER weight sync (otherwise repeated
# bare onload/offload cycles corrupt CUDA graph state).
if rollout_id >= args.num_critic_only_steps:
if args.offload_rollout:
await rollout_manager.onload_weights.remote()
await actor_model.update_weights()
if args.offload_rollout:
await rollout_manager.onload_kv.remote()
elif args.offload_rollout and args.colocate:
# Critic-only + colocate: engines were offloaded for critic training,
# must onload with proper weight sync before next generate().
await rollout_manager.onload_weights.remote()
await actor_model.update_weights()
await rollout_manager.onload_kv.remote()
# Skip eval during critic-only steps — actor weights (and thus
# rollout engine weights) are unchanged, so eval results would be
# identical to the previous eval.
if rollout_id >= args.num_critic_only_steps:
if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch):
await rollout_manager.eval.remote(rollout_id)
await rollout_manager.dispose.remote()
if __name__ == "__main__":
args = parse_args()
try:
asyncio.run(train(args))
finally:
finish_tracking()