mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Resolves op35-150 #1. Every driver released its objects on its last lines, so a raise from training, evaluation, or any of the dispose calls left the Ray CommandActor without a shutdown request and the run without its process trees reaped. Each driver now takes a disposer its entry point opens for it and registers each object's teardown where the object is created, so an exit path nobody thought of still releases everything in reverse creation order. `with_exit_stack` is what opens it: `asyncio.run(with_exit_stack(train, args))` replaces both the bare `asyncio.run` and the `finally: finish_tracking()` that only three of the four drivers had. The two teardowns every driver shares are no longer every driver's business. `init_orchestration_script` is what starts the tracking and launches the worker manager, so it takes the disposer too and registers `finish_tracking` and `shutdown_worker_manager` itself; no driver names either any more. The async driver's trailing `await eval_dispatcher.drain()` becomes a registration made where the dispatcher is created. That is after the executor and the models are registered, so the drain runs before them on the way out and an in-flight eval still finds the engines and the model it is reading. The drain any raise above it used to skip now happens on every exit path. A teardown that fails no longer stops the ones behind it: the stack runs every callback and chains what they raise onto the failure that started the unwind. The AST test that guards this selected no script at all, because no driver has named `launch_worker_manager` since `init_orchestration_script` took the launch over. It selects on that name now, and asserts each driver hands its disposer to the shared composition root and is entered through `with_exit_stack`.
134 lines
6.1 KiB
Python
134 lines
6.1 KiB
Python
import asyncio
|
|
import logging
|
|
import os
|
|
from contextlib import AsyncExitStack
|
|
|
|
from miles.ray.placement_group import (
|
|
create_rollout_components,
|
|
create_training_models,
|
|
maybe_start_api_server,
|
|
update_weights,
|
|
)
|
|
from miles.ray.rollout.eval_dispatch import EvalDispatcher
|
|
from miles.utils.arguments import parse_args, validate_async_off_policy_correction
|
|
from miles.utils.async_utils import eager_create_task, with_exit_stack
|
|
from miles.utils.data import remove_rollout_data_refs, remove_train_output_refs
|
|
from miles.utils.ft_utils.mini_ft_controller import maybe_start_mini_ft_controller
|
|
from miles.utils.misc import should_run_periodic_action
|
|
from miles.utils.orchestration_utils import init_orchestration_script
|
|
|
|
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, *, disposer: AsyncExitStack):
|
|
assert not args.colocate, "Colocation is not supported for async training."
|
|
validate_async_off_policy_correction(args)
|
|
init_orchestration_script(args, disposer=disposer)
|
|
|
|
# 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)
|
|
disposer.push_async_callback(inference_controller.dispose)
|
|
disposer.push_async_callback(rollout_executor.dispose)
|
|
|
|
# create the actor and critic models
|
|
actor_model, critic_model = await create_training_models(args, rollout_executor)
|
|
if critic_model is not None:
|
|
disposer.push_async_callback(critic_model.dispose)
|
|
disposer.push_async_callback(actor_model.dispose)
|
|
|
|
maybe_start_api_server(args, trainer_models={"actor": actor_model}, inference_controller=inference_controller)
|
|
maybe_start_mini_ft_controller(args)
|
|
|
|
# always update weight first so that sglang has the loaded weights from training.
|
|
await update_weights(args, actor_model, rollout_executor, inference_controller)
|
|
|
|
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,
|
|
)
|
|
|
|
eval_dispatcher = EvalDispatcher(args, actor_model, rollout_executor)
|
|
disposer.push_async_callback(eval_dispatcher.drain)
|
|
|
|
if args.eval_interval is not None and args.start_rollout_id == 0 and not args.skip_eval_before_train:
|
|
await inference_controller.prepare_eval()
|
|
await eval_dispatcher.dispatch(0, hf_dir=args.hf_checkpoint)
|
|
|
|
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 def prepare_and_generate(rollout_id):
|
|
await inference_controller.prepare_rollout(rollout_id)
|
|
return await rollout_executor.get(rollout_id)
|
|
|
|
# async train loop.
|
|
rollout_data_next_future = await eager_create_task(prepare_and_generate(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 = await eager_create_task(prepare_and_generate(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()
|
|
remove_train_output_refs(values)
|
|
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_executor.save(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 update_weights(args, actor_model, rollout_executor, inference_controller, rollout_id=rollout_id)
|
|
|
|
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
|
|
|
|
|
|
if __name__ == "__main__":
|
|
args = parse_args()
|
|
asyncio.run(with_exit_stack(train, args))
|