import asyncio import copy import itertools import logging import os from argparse import Namespace from pathlib import Path from miles.backends.megatron_utils.megatron_config import resolve_megatron_config from miles.ray.placement_group import create_rollout_components, maybe_start_api_server, update_weights from miles.ray.specs.train import compute_trainer_configs from miles.utils.arguments import parse_args from miles.utils.async_utils import Disposer, wait_cancelling_pending_on_first_completion, with_disposer from miles.utils.data import remove_rollout_data_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.multi_policy.checkpoint_state import MultiPolicyCheckpointState from miles.utils.multi_policy.parker import Parker from miles.utils.multi_policy.utils import ( TrainerInfo, assert_consistent_restore, create_trainers, define_policy_metric_groups, validate_multi_policy_args, ) from miles.utils.orchestration_utils import init_orchestration_script from miles.utils.workers.worker_handle import BaseWorkerHandle logger = logging.getLogger(__name__) async def train_multi_policy(args, *, disposer: Disposer) -> None: megatron_config = resolve_megatron_config(args) validate_multi_policy_args(args, megatron_config=megatron_config) _worker_manager = init_orchestration_script(args, disposer=disposer) define_policy_metric_groups(megatron_config) inference_controller, rollout_executor, num_rollout_per_epoch = await create_rollout_components(args) disposer.add(inference_controller, rollout_executor) trainers = await create_trainers(args, rollout_executor=rollout_executor) for trainer in trainers.values(): disposer.add(trainer.handle) assert_consistent_restore(args, trainers=trainers, leader_model_id=megatron_config.leader_model_id) maybe_start_api_server( args, trainer_models={ trainer_config.trainer_id: trainers[trainer_config.model_id].handle for trainer_config in compute_trainer_configs(args) }, inference_controller=inference_controller, ) maybe_start_mini_ft_controller(args) for model_id, trainer in trainers.items(): await update_weights( _startup_args(args, trainer=trainer), trainer.handle, rollout_executor, inference_controller, trainer_model_id=model_id, ) 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, model_id=model_id, ) leader_start_rollout_id = trainers[megatron_config.leader_model_id].start_rollout_id if args.eval_interval is not None and leader_start_rollout_id == 0 and not args.skip_eval_before_train: await inference_controller.prepare_eval() await rollout_executor.eval(0) save_parker = Parker(num_followers=len(trainers) - 1) eval_parker = Parker(num_followers=len(trainers) - 1) run_ended = asyncio.Event() rollout_ids: dict[str, int] = {} tasks = [ asyncio.create_task( _run_policy( args, trainer=trainer, is_leader=trainer.model_id == megatron_config.leader_model_id, trainers=trainers, inference_controller=inference_controller, rollout_executor=rollout_executor, save_parker=save_parker, eval_parker=eval_parker, run_ended=run_ended, rollout_ids=rollout_ids, num_rollout_per_epoch=num_rollout_per_epoch, ) ) for trainer in trainers.values() ] await wait_cancelling_pending_on_first_completion(tasks, on_first_completion=run_ended.set) def _startup_args(args, *, trainer: TrainerInfo) -> Namespace: ans = copy.copy(args) ans.start_rollout_id = trainer.start_rollout_id return ans async def _run_policy( args, *, trainer: TrainerInfo, is_leader: bool, run_ended: asyncio.Event, trainers: dict[str, TrainerInfo], inference_controller: BaseWorkerHandle, rollout_executor: BaseWorkerHandle, save_parker: Parker, eval_parker: Parker, rollout_ids: dict[str, int], num_rollout_per_epoch: int | None, ) -> None: model_id = trainer.model_id rollout_ids_iter = ( range(trainer.start_rollout_id, args.num_rollout) if is_leader else itertools.count(trainer.start_rollout_id) ) for rollout_id in rollout_ids_iter: if run_ended.is_set(): return rollout_ids[model_id] = rollout_id await inference_controller.prepare_rollout(rollout_id, model_id=model_id) rollout_data_pack = await rollout_executor.get(rollout_id, trainer_model_id=model_id) await trainer.handle.train(rollout_id, rollout_data_pack) remove_rollout_data_refs(args, rollout_data_pack) if is_leader: await _maybe_save_globally( args, model_id=model_id, trainers=trainers, rollout_executor=rollout_executor, parker=save_parker, rollout_ids=rollout_ids, rollout_id=rollout_id, num_rollout_per_epoch=num_rollout_per_epoch, ) else: await save_parker.maybe_park_follower() if (rollout_id + 1) % args.update_weights_interval == 0: await update_weights( args, trainer.handle, rollout_executor, inference_controller, rollout_id=rollout_id, trainer_model_id=model_id, ) if is_leader: if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch, args.num_rollout): async with eval_parker.with_all_parked(): await inference_controller.prepare_eval() await rollout_executor.eval(rollout_id) else: await eval_parker.maybe_park_follower() if ( is_leader and (x := args.debug_exit_after_rollout) is not None and (rollout_id - trainer.start_rollout_id + 1) >= x ): logger.info(f"debug_exit_after_rollout={x} reached at rollout_id={rollout_id}, exiting") break async def _maybe_save_globally( args, *, model_id: str, trainers: dict[str, TrainerInfo], rollout_executor: BaseWorkerHandle, parker: Parker, rollout_ids: dict[str, int], rollout_id: int, num_rollout_per_epoch: int | None, ) -> None: external_save = args.save_trigger_sentinel is not None and os.path.exists(args.save_trigger_sentinel) if not external_save and not should_run_periodic_action( rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout ): return async with parker.with_all_parked(): await asyncio.gather( *( trainer.handle.save_model(rollout_ids[trainer_model_id], force_sync=True) for trainer_model_id, trainer in trainers.items() ) ) await rollout_executor.save(rollout_id) if args.save is not None: MultiPolicyCheckpointState(leader_model_id=model_id, rollout_ids=dict(rollout_ids)).save(Path(args.save)) if external_save: os.remove(args.save_trigger_sentinel) if __name__ == "__main__": args = parse_args() asyncio.run(with_disposer(train_multi_policy, args))