config: delete dead ensure_model_parallel_initialized (#39137)

This commit is contained in:
Cheng Wan
2026-09-12 17:27:12 -07:00
committed by GitHub
parent 6804eeaabe
commit fa260f26da
@@ -11,8 +11,7 @@ It takes over the control of the distributed environment from PyTorch.
The typical workflow is:
- call `init_distributed_environment` to initialize the distributed environment.
- call `initialize_model_parallel` or `ensure_model_parallel_initialized` to
initialize the model parallel groups.
- call `initialize_model_parallel` to initialize the model parallel groups.
- any code dealing with the distributed stuff
@@ -2885,46 +2884,6 @@ def create_custom_parallel_group(
return my_new_group
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
expert_model_parallel_size: int,
pipeline_model_parallel_size: int,
decode_context_parallel_size: int = 1,
backend: Optional[str] = None,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
or ensure tensor-parallel and pipeline-parallel sizes are equal to expected
values if the model parallel groups are initialized.
"""
backend = backend or torch.distributed.get_backend(get_world_group().device_group)
if not model_parallel_is_initialized():
initialize_model_parallel(
tensor_model_parallel_size=tensor_model_parallel_size,
expert_model_parallel_size=expert_model_parallel_size,
pipeline_model_parallel_size=pipeline_model_parallel_size,
decode_context_parallel_size=decode_context_parallel_size,
backend=backend,
)
return
assert get_tensor_model_parallel_world_size() == tensor_model_parallel_size, (
"tensor parallel group already initialized, but of unexpected size: "
f"{get_tensor_model_parallel_world_size()=} vs. "
f"{tensor_model_parallel_size=}"
)
pp_world_size = get_pp_group().world_size
assert pp_world_size == pipeline_model_parallel_size, (
"pipeline parallel group already initialized, but of unexpected size: "
f"{pp_world_size=} vs. "
f"{pipeline_model_parallel_size=}"
)
if decode_context_parallel_size > 1:
dcp_world_size = get_dcp_group().world_size
assert dcp_world_size == decode_context_parallel_size, (
f"decode context parallel group already initialized, but of unexpected size: {dcp_world_size=} {decode_context_parallel_size=}"
)
def model_parallel_is_initialized():
"""Check if tensor and pipeline parallel groups are initialized."""
return _TP is not None and _PP is not None