Files
miles/train.py
T

195 lines
8.2 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, 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 rollout_id + 1 < args.num_rollout or should_run_periodic_action(
rollout_id, args.eval_interval, num_rollout_per_epoch
):
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()