[Mooncake] Fix ProcessGroup API imports (#32403)
This commit is contained in:
@@ -311,7 +311,7 @@ class GroupCoordinator:
|
||||
for ranks in group_ranks:
|
||||
subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT
|
||||
if "mooncake" in torch_distributed_backend:
|
||||
from mooncake.ep import MooncakeBackendOptions
|
||||
from mooncake.pg import MooncakeBackendOptions
|
||||
|
||||
pg_active_size = len(ranks)
|
||||
if not recovered_rank and max_world_size is not None:
|
||||
@@ -2114,7 +2114,7 @@ def init_distributed_environment(
|
||||
_MODEL_PARALLEL_GROUP_TIMEOUT = timeout
|
||||
|
||||
if backend == "mooncake":
|
||||
from mooncake.ep import MooncakeBackendOptions
|
||||
from mooncake.pg import MooncakeBackendOptions
|
||||
|
||||
use_max_ws = max_world_size and max_world_size > world_size
|
||||
ar_size = max_world_size if use_max_ws else world_size
|
||||
|
||||
@@ -351,8 +351,10 @@ def _map_global_to_group_local_ranks(
|
||||
return [rank_to_local[rank] for rank in global_ranks if rank in rank_to_local]
|
||||
|
||||
|
||||
def _wait_for_peer_state(mooncake_ep, backend, ranks: List[int]) -> None:
|
||||
while not all(mooncake_ep.get_peer_state(backend, ranks)):
|
||||
def _wait_for_peer_state(backend, ranks: List[int]) -> None:
|
||||
from mooncake.pg import get_peer_state
|
||||
|
||||
while not all(get_peer_state(backend, ranks)):
|
||||
time.sleep(_PEER_STATE_POLL_INTERVAL_SEC)
|
||||
|
||||
|
||||
@@ -368,13 +370,13 @@ def _maybe_create_message_queue(group) -> None:
|
||||
|
||||
|
||||
def _try_recover_world(global_ranks: List[int]) -> bool:
|
||||
from mooncake import ep as mooncake_ep
|
||||
from mooncake.pg import get_peer_state, recover_ranks
|
||||
|
||||
world_backend = torch.distributed.group.WORLD
|
||||
if not all(mooncake_ep.get_peer_state(world_backend, global_ranks)):
|
||||
if not all(get_peer_state(world_backend, global_ranks)):
|
||||
return False
|
||||
|
||||
mooncake_ep.recover_ranks(world_backend, global_ranks)
|
||||
recover_ranks(world_backend, global_ranks)
|
||||
logger.debug("[Elastic EP][recover] WORLD recover_ranks(%s) done", global_ranks)
|
||||
return True
|
||||
|
||||
@@ -393,17 +395,17 @@ def try_recover_ranks(global_ranks: List[int]) -> bool:
|
||||
if not _try_recover_world(global_ranks):
|
||||
return False
|
||||
|
||||
from mooncake import ep as mooncake_ep
|
||||
from mooncake.pg import recover_ranks
|
||||
|
||||
for group in _iter_live_parallel_groups():
|
||||
local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks)
|
||||
if not local_ranks:
|
||||
continue
|
||||
|
||||
_wait_for_peer_state(mooncake_ep, group.device_group, local_ranks)
|
||||
mooncake_ep.recover_ranks(group.device_group, local_ranks)
|
||||
_wait_for_peer_state(mooncake_ep, group.cpu_group, local_ranks)
|
||||
mooncake_ep.recover_ranks(group.cpu_group, local_ranks)
|
||||
_wait_for_peer_state(group.device_group, local_ranks)
|
||||
recover_ranks(group.device_group, local_ranks)
|
||||
_wait_for_peer_state(group.cpu_group, local_ranks)
|
||||
recover_ranks(group.cpu_group, local_ranks)
|
||||
_maybe_create_message_queue(group)
|
||||
|
||||
_refresh_ep_members()
|
||||
@@ -411,9 +413,9 @@ def try_recover_ranks(global_ranks: List[int]) -> bool:
|
||||
|
||||
|
||||
def _join_world_group() -> None:
|
||||
from mooncake import ep as mooncake_ep
|
||||
from mooncake.pg import join_group
|
||||
|
||||
mooncake_ep.join_group(torch.distributed.group.WORLD)
|
||||
join_group(torch.distributed.group.WORLD)
|
||||
|
||||
|
||||
def join_scale_process_group() -> None:
|
||||
@@ -424,14 +426,14 @@ def join_scale_process_group() -> None:
|
||||
|
||||
def join_process_groups() -> None:
|
||||
"""Rejoin WORLD and every launch-time parallel group after recovery."""
|
||||
from mooncake import ep as mooncake_ep
|
||||
from mooncake.pg import join_group
|
||||
|
||||
_join_world_group()
|
||||
for group in _iter_live_parallel_groups():
|
||||
if group.world_size <= 1:
|
||||
continue
|
||||
mooncake_ep.join_group(group.device_group)
|
||||
mooncake_ep.join_group(group.cpu_group)
|
||||
join_group(group.device_group)
|
||||
join_group(group.cpu_group)
|
||||
_maybe_create_message_queue(group)
|
||||
|
||||
_refresh_ep_members()
|
||||
|
||||
Reference in New Issue
Block a user