Files
miles/train_async.py
T
Tom Chen e2bb11a9d1 Own every driver's teardown in a disposer
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_disposer` is what opens it:
`asyncio.run(with_disposer(train, args))` replaces both the bare `asyncio.run`
and the `finally: finish_tracking()` that only three of the four drivers had.

The disposer is what a driver registers against, and it takes the objects
themselves: `disposer.add(inference_controller, rollout_executor)` is the whole
registration, because an object carrying a `dispose` is registered by that
method and anything else has to be callable. A callable it can await is pushed
as an async callback and a plain one as a callback, so a teardown needing an
argument is handed over as a `functools.partial` and is still awaited. An object
a run never created is skipped rather than refused, so a driver whose critic is
optional writes `disposer.add(critic_model, actor_model)` and keeps the order
its old `finally` released them in. Naming the objects rather than their
teardowns is what keeps a driver from registering the wrong half of a pair it
just created.

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. It still
hands the manager back and every driver still binds it, because a Ray actor
handle nothing holds is destroyed the moment it is dropped.

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_disposer`.
2026-09-04 10:13:33 +08:00

130 lines
5.9 KiB
Python

import asyncio
import logging
import os
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 Disposer, eager_create_task, with_disposer
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: Disposer):
assert not args.colocate, "Colocation is not supported for async training."
validate_async_off_policy_correction(args)
_worker_manager = 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.add(inference_controller, rollout_executor)
# create the actor and critic models
actor_model, critic_model = await create_training_models(args, rollout_executor)
disposer.add(critic_model, actor_model)
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.add(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_disposer(train, args))