forked from radixark/miles
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_async.py
More file actions
134 lines (112 loc) · 5.8 KB
/
Copy pathtrain_async.py
File metadata and controls
134 lines (112 loc) · 5.8 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
import asyncio
import logging
import os
from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models
from miles.utils import object_store
from miles.utils.arguments import parse_args, validate_async_off_policy_correction
from miles.utils.audit_utils.process_identity import MainProcessIdentity
from miles.utils.data import remove_rollout_data_refs
from miles.utils.debug_utils.periodic_py_spy import maybe_start_periodic_pyspy_dump
from miles.utils.ft_utils.control_server.server import start_control_server
from miles.utils.ft_utils.mini_ft_controller import maybe_start_mini_ft_controller
from miles.utils.logging_utils import configure_logger
from miles.utils.misc import should_run_periodic_action
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
logger = logging.getLogger(__name__)
# The framework supports other asynchronous approaches such as fully async (see miles/rollout/fully_async_rollout.py).
async def train(args):
if getattr(args, "rollout_only_from_checkpoint", False):
raise ValueError("--rollout-only-from-checkpoint is supported only by synchronous train.py")
assert not args.colocate, "Colocation is not supported for async training."
validate_async_off_policy_correction(args)
configure_logger(args, source=MainProcessIdentity())
maybe_start_periodic_pyspy_dump()
# allocate the GPUs
pgs = create_placement_groups(args)
object_store.init_instance(args, contribute_segment=False)
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.control_server_port:
start_control_server(
actor_model=actor_model,
rollout_manager=rollout_manager,
port=args.control_server_port,
ft_components=args.ft_components,
)
maybe_start_mini_ft_controller(args)
# 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",
allow_quant_error=args.check_weight_update_allow_quant_error,
selector=args.check_weight_update_selector,
skip_list=args.check_weight_update_skip_list,
)
if args.eval_interval is not None and args.start_rollout_id == 0 and not args.skip_eval_before_train:
await rollout_manager.eval.remote(0)
async def save_training_model(model, rollout_id, force_sync):
if args.use_critic and args.offload_train:
await model.onload()
await model.save_model(rollout_id, force_sync=force_sync)
if args.use_critic and args.offload_train:
await model.offload()
# async train loop.
rollout_data_next_future = rollout_manager.generate.remote(args.start_rollout_id)
for rollout_id in range(args.start_rollout_id, args.num_rollout):
# Sync the last generation
if rollout_data_next_future is not None:
rollout_data_curr_ref = await rollout_data_next_future
# Start the next rollout early.
if rollout_id + 1 < args.num_rollout:
rollout_data_next_future = rollout_manager.generate.remote(rollout_id + 1)
if args.use_critic:
values = await critic_model.train(rollout_id, rollout_data_curr_ref)
if args.offload_train:
await critic_model.offload()
if rollout_id >= args.num_critic_only_steps:
await actor_model.train(rollout_id, rollout_data_curr_ref, external_data=values)
if args.offload_train:
await actor_model.offload()
else:
await actor_model.train(rollout_id, rollout_data_curr_ref)
remove_rollout_data_refs(args, rollout_data_curr_ref)
external_save = args.save_trigger_sentinel is not None and os.path.exists(args.save_trigger_sentinel)
if external_save or should_run_periodic_action(
rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout
):
force_sync = external_save or rollout_id == args.num_rollout - 1
await save_training_model(actor_model, rollout_id, force_sync)
if args.use_critic:
await save_training_model(critic_model, rollout_id, force_sync)
await rollout_manager.save.remote(rollout_id)
if external_save:
os.remove(args.save_trigger_sentinel)
if (rollout_id + 1) % args.update_weights_interval == 0:
# sync generate before update weights to prevent update weight in the middle of generation
rollout_data_curr_ref = (await x) if (x := rollout_data_next_future) is not None else None
rollout_data_next_future = None
await actor_model.update_weights(rollout_id=rollout_id)
if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch):
await rollout_manager.eval.remote(rollout_id)
if (
args.debug_exit_after_rollout is not None
and (rollout_id - args.start_rollout_id + 1) >= args.debug_exit_after_rollout
):
logger.info(
"debug_exit_after_rollout=%d reached at rollout_id=%d, exiting",
args.debug_exit_after_rollout,
rollout_id,
)
break
await rollout_manager.dispose.remote()
if __name__ == "__main__":
args = parse_args()
try:
asyncio.run(train(args))
finally:
finish_tracking()