Fix mooncake scale joiner groups (#38503)
Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
co-authored by
Po-Han Huang
parent
fb01a079a6
commit
1895cabfa8
@@ -293,9 +293,11 @@ def _init_parallel_groups(
|
||||
decode_context_parallel_size=dcp_size,
|
||||
duplicate_tp_group=get_disagg().enable_pdmux,
|
||||
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||
recovered_rank=is_ep_joiner,
|
||||
# Only WORLD is extended during scale-up. The joiner's model-parallel
|
||||
# groups are fixed groups local to its launch cohort.
|
||||
recovered_rank=is_ep_joiner and not is_scale_joiner,
|
||||
rank_offset=rank_offset,
|
||||
max_world_size=get_parallel().max_ep_size,
|
||||
max_world_size=None if is_scale_joiner else get_parallel().max_ep_size,
|
||||
)
|
||||
_tag_groups_for_flashinfer_allreduce_only()
|
||||
initialize_dp_attention(
|
||||
|
||||
@@ -313,6 +313,8 @@ class GroupCoordinator:
|
||||
# by _tag_groups_for_flashinfer_allreduce_only() after group init.
|
||||
self._fi_workspace_hint: Optional[str] = None
|
||||
self.local_size = get_int_env_var("LOCAL_SIZE", 0)
|
||||
# Scale cohorts create these groups without the serving ranks.
|
||||
use_local_synchronization = rank_offset > 0 and not recovered_rank
|
||||
|
||||
if is_cuda_alike():
|
||||
device_id = (
|
||||
@@ -335,7 +337,7 @@ class GroupCoordinator:
|
||||
from mooncake.pg import MooncakeBackendOptions
|
||||
|
||||
pg_active_size = len(ranks)
|
||||
if not recovered_rank and max_world_size is not None:
|
||||
if max_world_size is not None:
|
||||
assert max_world_size >= len(ranks), (
|
||||
f"max_world_size ({max_world_size}) must be >= "
|
||||
f"group size ({len(ranks)})"
|
||||
@@ -349,7 +351,7 @@ class GroupCoordinator:
|
||||
pg_active_ranks_cpu = torch.zeros(pg_active_size, dtype=torch.int32)
|
||||
pg_active_ranks_cpu[: len(ranks)] = 1
|
||||
|
||||
if not recovered_rank and max_world_size is not None:
|
||||
if max_world_size is not None:
|
||||
dev_opts = MooncakeBackendOptions(
|
||||
pg_active_ranks, recovered_rank, max_world_size
|
||||
)
|
||||
@@ -370,6 +372,7 @@ class GroupCoordinator:
|
||||
pg_options=dev_opts,
|
||||
timeout=subgroup_timeout,
|
||||
group_desc=f"{group_name}:device",
|
||||
use_local_synchronization=use_local_synchronization,
|
||||
)
|
||||
cpu_group = torch.distributed.new_group(
|
||||
ranks,
|
||||
@@ -377,6 +380,7 @@ class GroupCoordinator:
|
||||
pg_options=cpu_opts,
|
||||
timeout=subgroup_timeout,
|
||||
group_desc=f"{group_name}:cpu",
|
||||
use_local_synchronization=use_local_synchronization,
|
||||
)
|
||||
else:
|
||||
active_ranks = torch.ones(
|
||||
@@ -390,6 +394,7 @@ class GroupCoordinator:
|
||||
pg_options=pg_options,
|
||||
timeout=subgroup_timeout,
|
||||
group_desc=f"{group_name}:device",
|
||||
use_local_synchronization=use_local_synchronization,
|
||||
)
|
||||
# a group with `gloo` backend, to allow direct coordination
|
||||
# between processes through the CPU.
|
||||
@@ -398,6 +403,7 @@ class GroupCoordinator:
|
||||
backend="gloo",
|
||||
timeout=gloo_timeout,
|
||||
group_desc=f"{group_name}:cpu",
|
||||
use_local_synchronization=use_local_synchronization,
|
||||
)
|
||||
if self.rank in ranks:
|
||||
self.ranks = ranks
|
||||
@@ -2028,7 +2034,11 @@ def get_world_group() -> GroupCoordinator:
|
||||
|
||||
|
||||
def init_world_group(
|
||||
ranks: List[int], local_rank: int, backend: str, recovered_rank: bool = False
|
||||
ranks: List[int],
|
||||
local_rank: int,
|
||||
backend: str,
|
||||
recovered_rank: bool = False,
|
||||
max_world_size: Optional[int] = None,
|
||||
) -> GroupCoordinator:
|
||||
return GroupCoordinator(
|
||||
group_ranks=[ranks],
|
||||
@@ -2043,6 +2053,7 @@ def init_world_group(
|
||||
use_npu_communicator=False,
|
||||
group_name="world",
|
||||
recovered_rank=recovered_rank,
|
||||
max_world_size=max_world_size,
|
||||
)
|
||||
|
||||
|
||||
@@ -2441,7 +2452,11 @@ def init_distributed_environment(
|
||||
if _WORLD is None:
|
||||
ranks = list(range(torch.distributed.get_world_size()))
|
||||
_WORLD = init_world_group(
|
||||
ranks, local_rank, backend, recovered_rank=recovered_rank
|
||||
ranks=ranks,
|
||||
local_rank=local_rank,
|
||||
backend=backend,
|
||||
recovered_rank=recovered_rank,
|
||||
max_world_size=max_world_size,
|
||||
)
|
||||
else:
|
||||
assert _WORLD.world_size == torch.distributed.get_world_size(), (
|
||||
@@ -2524,7 +2539,7 @@ def initialize_model_parallel(
|
||||
# Joiners construct their local TP/PP layout in global rank space.
|
||||
world_size: int = (
|
||||
tensor_model_parallel_size * pipeline_model_parallel_size
|
||||
if recovered_rank
|
||||
if recovered_rank or rank_offset > 0
|
||||
else torch.distributed.get_world_size()
|
||||
)
|
||||
|
||||
@@ -2613,6 +2628,8 @@ def initialize_model_parallel(
|
||||
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
|
||||
group_name="dcp",
|
||||
recovered_rank=recovered_rank,
|
||||
rank_offset=rank_offset,
|
||||
max_world_size=max_world_size,
|
||||
)
|
||||
if get_tensor_model_parallel_rank() == 0:
|
||||
logger.info(
|
||||
@@ -2830,6 +2847,9 @@ def initialize_model_parallel(
|
||||
backend,
|
||||
use_custom_allreduce=False,
|
||||
group_name="self_pp",
|
||||
recovered_rank=recovered_rank,
|
||||
rank_offset=rank_offset,
|
||||
max_world_size=max_world_size,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user