mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
195 lines
8.2 KiB
Python
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()
|