mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Drop the drivers' unused copy of the worker manager handle
init_orchestration_script registers the worker manager's shutdown with the disposer, and that registration already keeps the actor handle alive for the whole run. The handle it also returned was bound to an unused variable in every driver, so the return value goes away.
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
from argparse import Namespace
|
||||
from functools import partial
|
||||
|
||||
from ray.actor import ActorHandle
|
||||
|
||||
from miles.ray.wiring import launch_worker_manager, shutdown_worker_manager
|
||||
from miles.utils import object_store
|
||||
from miles.utils.async_utils import Disposer
|
||||
@@ -16,7 +14,7 @@ from miles.utils.test_utils.fault_injector.models import FaultHookOwner
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
|
||||
|
||||
|
||||
def init_orchestration_script(args: Namespace, *, disposer: Disposer) -> ActorHandle | None:
|
||||
def init_orchestration_script(args: Namespace, *, disposer: Disposer) -> None:
|
||||
event_logger_checkpoint.restore(args)
|
||||
configure_logger(args, source=SimpleProcessIdentity(component="main"))
|
||||
maybe_start_periodic_pyspy_dump()
|
||||
@@ -26,4 +24,3 @@ def init_orchestration_script(args: Namespace, *, disposer: Disposer) -> ActorHa
|
||||
disposer.add(partial(shutdown_worker_manager, worker_manager))
|
||||
object_store.init_instance(args, contribute_segment=False)
|
||||
fault_hook_controller.configure(resources=FaultHookResources(args=args), owner=FaultHookOwner.ORCHESTRATOR)
|
||||
return worker_manager
|
||||
|
||||
+1
-1
@@ -37,7 +37,7 @@ async def serve(args, *, disposer: Disposer):
|
||||
trainer_token_limit = args.max_tokens_per_gpu // pad_size * pad_size
|
||||
max_tokens_per_datum = min(max_tokens_per_datum, trainer_token_limit)
|
||||
assert max_tokens_per_datum > 0, "trainer token budget must fit at least one padding block"
|
||||
_worker_manager = init_orchestration_script(args, disposer=disposer)
|
||||
init_orchestration_script(args, disposer=disposer)
|
||||
|
||||
capability = get_backend_capability(args)
|
||||
await resolve_router_addrs(args, router_providers=compute_router_providers(args, capability=capability))
|
||||
|
||||
@@ -20,9 +20,6 @@ RAY_USING_MODULES = {
|
||||
"miles/ray/placement_group.py": "launcher closure: placement groups are how ray is asked to schedule",
|
||||
"miles/ray/wiring.py": "launcher closure: driver shutdown kills the manager it launched by ActorHandle",
|
||||
"miles/utils/ray_utils.py": "launcher closure: node lookup and pinning options for the launcher's own calls",
|
||||
"miles/utils/orchestration_utils.py": (
|
||||
"launcher closure: the shared driver composition root returns its ray worker manager"
|
||||
),
|
||||
"miles/utils/workers/ray_worker_manager.py": "launcher closure: it is the launcher",
|
||||
"miles/utils/workers/ray_worker_handle.py": "launcher closure: the handle of the ray communication mode itself",
|
||||
"miles/utils/workers/worker_provider/ray.py": "launcher closure: it reads the launcher's own bookkeeping",
|
||||
|
||||
@@ -24,7 +24,7 @@ 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"
|
||||
_worker_manager = init_orchestration_script(args, disposer=disposer)
|
||||
init_orchestration_script(args, disposer=disposer)
|
||||
|
||||
if args.colocate_memory_peak_device == "gpu":
|
||||
assert (
|
||||
|
||||
+1
-1
@@ -23,7 +23,7 @@ logger = logging.getLogger(__name__)
|
||||
async def train(args, *, disposer: Disposer):
|
||||
assert not args.colocate or args.fully_async, "Colocation is only supported for async training with --fully-async."
|
||||
validate_async_off_policy_correction(args)
|
||||
_worker_manager = init_orchestration_script(args, disposer=disposer)
|
||||
init_orchestration_script(args, disposer=disposer)
|
||||
|
||||
# create the rollout manager, with sglang engines inside.
|
||||
# need to initialize rollout manager first to calculate num_rollout
|
||||
|
||||
@@ -31,7 +31,7 @@ logger = logging.getLogger(__name__)
|
||||
async def train_multi_policy(args, *, disposer: Disposer) -> None:
|
||||
megatron_config = args.raw_megatron
|
||||
validate_multi_policy_args(args, megatron_config=megatron_config)
|
||||
_worker_manager = init_orchestration_script(args, disposer=disposer)
|
||||
init_orchestration_script(args, disposer=disposer)
|
||||
|
||||
define_policy_metric_groups(megatron_config)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user