Files
miles/train.py
T
Tom Chen 8e8137cbda 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. 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.

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:02:22 +08:00

149 lines
6.1 KiB
Python

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,
maybe_start_api_server,
update_weights,
)
from miles.utils.arguments import parse_args
from miles.utils.async_utils import Disposer, 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__)
async def train(args, *, disposer: Disposer):
assert not args.fully_async, "--fully-async requires the async driver: run train_async.py"
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)
if critic_model is not None:
disposer.add(critic_model)
disposer.add(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,
)
if args.offload_rollout:
await inference_controller.onload_kv()
# special case for eval-only
if args.num_rollout == 0 and args.eval_interval is not None:
await inference_controller.prepare_eval()
await rollout_executor.eval(rollout_id=0)
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(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 == args.start_rollout_id and not args.skip_eval_before_train:
await inference_controller.prepare_eval()
await rollout_executor.eval(rollout_id)
await inference_controller.prepare_rollout(rollout_id)
rollout_data_pack = await rollout_executor.get(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 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()
remove_train_output_refs(values)
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)
await offload_train()
if args.offload_rollout:
await inference_controller.onload_weights()
await update_weights(args, actor_model, rollout_executor, inference_controller, 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):
await inference_controller.prepare_eval()
await rollout_executor.eval(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
if __name__ == "__main__":
args = parse_args()
asyncio.run(with_disposer(train, args))