Repository navigation
Expand file tree
/
Copy pathtrain.py
More file actions
191 lines (161 loc) · 7.95 KB
/
Copy pathtrain.py
File metadata and controls
191 lines (161 loc) · 7.95 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
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
import asyncio
import logging
import os
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_rollout_components, create_training_models, update_weights
from miles.ray.rollout.eval_dispatch import EvalDispatcher
from miles.ray.wiring import launch_worker_manager
from miles.utils import object_store
from miles.utils.arguments import parse_args
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.api_server.server import start_api_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.lora import lora_rollout_enabled
from miles.utils.misc import should_run_periodic_action
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
logger = logging.getLogger(__name__)
async def train(args):
assert not args.fully_async, "--fully-async requires the async driver: run train_async.py"
configure_logger(args, source=MainProcessIdentity())
maybe_start_periodic_pyspy_dump()
_worker_manager = launch_worker_manager(args)
object_store.init_instance(args, contribute_segment=False)
init_tracking(args)
if args.colocate_memory_peak_device == "gpu":
assert (
args.offload_train and args.offload_rollout
), "--colocate-memory-peak-device gpu requires --offload-train and --offload-rollout"
assert not args.use_critic, "--colocate-memory-peak-device gpu is not wired for the critic path"
# create the rollout manager, with sglang engines inside.
# need to initialize rollout manager first to calculate num_rollout
inference_controller, rollout_executor, num_rollout_per_epoch = await create_rollout_components(args)
# create the actor and critic models
actor_model, critic_model = await create_training_models(args, inference_controller, rollout_executor)
if args.api_server_port:
start_api_server(
args=args,
actor_model=actor_model,
inference_controller=inference_controller,
host=args.api_server_host,
port=args.api_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 update_weights(actor_model, rollout_executor)
if args.check_weight_update_equal:
await inference_controller.check_weights(
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.offload_rollout:
await inference_controller.onload_kv()
eval_dispatcher = EvalDispatcher(args, actor_model, rollout_executor)
# special case for eval-only
if args.num_rollout == 0 and args.eval_interval is not None:
await inference_controller.prepare_eval()
await eval_dispatcher.dispatch(0, hf_dir=args.hf_checkpoint)
async def offload_train():
if args.use_critic:
return
if args.offload_train:
await actor_model.offload()
else:
await actor_model.clear_memory()
async def save(rollout_id, force_sync=False):
force_sync = force_sync or rollout_id == args.num_rollout - 1
async def save_training_model(model):
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()
if (not args.use_critic) or (rollout_id >= args.num_critic_only_steps):
await save_training_model(actor_model)
if args.use_critic:
await save_training_model(critic_model)
await rollout_executor.save.remote(rollout_id)
if args.num_rollout > args.start_rollout_id and args.eval_interval is not None and not args.skip_eval_before_train:
await inference_controller.prepare_eval()
if args.start_rollout_id == 0:
await eval_dispatcher.dispatch(0, hf_dir=args.hf_checkpoint)
else:
await eval_dispatcher.dispatch(args.start_rollout_id - 1)
# 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):
await inference_controller.prepare_rollout(rollout_id)
rollout_data_pack = await rollout_executor.get.remote(rollout_id)
if args.offload_rollout:
if args.colocate_memory_peak_device == "gpu":
await inference_controller.offload_kv()
await actor_model.onload()
await inference_controller.offload_weights()
else:
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 inference_controller.offload(tags=offload_tags)
if args.use_critic:
values = await critic_model.train(rollout_id, rollout_data_pack)
if args.offload_train:
await critic_model.offload()
if rollout_id >= args.num_critic_only_steps:
await actor_model.train(rollout_id, rollout_data_pack, external_data=values)
if args.offload_train:
await actor_model.offload()
else:
await actor_model.train(rollout_id, rollout_data_pack)
remove_rollout_data_refs(args, rollout_data_pack)
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
):
await save(rollout_id, force_sync=external_save)
if external_save:
os.remove(args.save_trigger_sentinel)
if args.colocate_memory_peak_device == "gpu":
await actor_model.clear_memory()
if lora_rollout_enabled(args):
await actor_model.offload_grad_buffer()
await inference_controller.onload_weights()
await offload_train()
else:
await offload_train()
if args.offload_rollout:
await inference_controller.onload_weights()
await update_weights(actor_model, rollout_executor, rollout_id=rollout_id)
if args.offload_rollout:
await inference_controller.onload_kv()
if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch, args.num_rollout):
await inference_controller.prepare_eval()
await eval_dispatcher.dispatch(rollout_id, force=rollout_id == args.num_rollout - 1)
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 eval_dispatcher.drain()
await rollout_executor.dispose.remote()
await inference_controller.dispose()
await actor_model.dispose()
if critic_model is not None:
await critic_model.dispose()
if __name__ == "__main__":
args = parse_args()
try:
asyncio.run(train(args))
finally:
finish_tracking()