From 5717834f1f3b803f79ca99a11148685294f3eacd Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 17 Mar 2026 21:21:42 +0800 Subject: [PATCH] [diffusion] refactor: cleanup parallel_state.py (#20760) Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- .../runtime/distributed/parallel_state.py | 453 +++--------------- 1 file changed, 58 insertions(+), 395 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 147cf1444..234ea66ec 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -59,14 +59,14 @@ from .group_coordinator import ( logger = init_logger(__name__) -_WORLD: Optional[GroupCoordinator] = None -_TP: Optional[GroupCoordinator] = None -_SP: Optional[SequenceParallelGroupCoordinator] = None -_PP: Optional[PipelineGroupCoordinator] = None -_CFG: Optional[GroupCoordinator] = None -_DP: Optional[GroupCoordinator] = None -_DIT: Optional[GroupCoordinator] = None -_VAE: Optional[GroupCoordinator] = None +_WORLD: GroupCoordinator | None = None +_TP: GroupCoordinator | None = None +_SP: SequenceParallelGroupCoordinator | None = None +_PP: PipelineGroupCoordinator | None = None +_CFG: GroupCoordinator | None = None +_DP: GroupCoordinator | None = None +_DIT: ProcessGroup | None = None +_VAE: ProcessGroup | None = None TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"]) @@ -116,10 +116,6 @@ def all_reduce_fake(tensor: torch.Tensor, group_name: str) -> torch.Tensor: return torch.empty_like(tensor) -_WORLD: GroupCoordinator | None = None -_NODE: GroupCoordinator | None = None - - def get_world_group() -> GroupCoordinator: assert _WORLD is not None, "world group is not initialized" return _WORLD @@ -137,7 +133,6 @@ def init_world_group( ) -# xDiT def init_parallel_group_coordinator( group_ranks: List[List[int]], local_rank: int, @@ -145,9 +140,7 @@ def init_parallel_group_coordinator( parallel_mode: str, **kwargs, ) -> GroupCoordinator: - """ - Returns a Group Coordinator for the given parallel mode - """ + """Return a group coordinator for the given parallel mode.""" assert parallel_mode in [ "data", "pipeline", @@ -180,39 +173,11 @@ def init_parallel_group_coordinator( ) -# def init_parallel_group_coordinator( -# group_ranks: list[list[int]], -# local_rank: int, -# backend: str, -# use_message_queue_broadcaster: bool = False, -# group_name: str | None = None, -# ) -> GroupCoordinator: -# return GroupCoordinator( -# group_ranks=group_ranks, -# local_rank=local_rank, -# torch_distributed_backend=backend, -# use_device_communicator=True, -# use_message_queue_broadcaster=use_message_queue_broadcaster, -# group_name=group_name, -# ) - - -_TP: GroupCoordinator | None = None - - def get_tp_group() -> GroupCoordinator: assert _TP is not None, "tensor model parallel group is not initialized" return _TP -_ENABLE_CUSTOM_ALL_REDUCE = True - - -def set_custom_all_reduce(enable: bool): - global _ENABLE_CUSTOM_ALL_REDUCE - _ENABLE_CUSTOM_ALL_REDUCE = enable - - def init_distributed_environment( world_size: int = 1, rank: int = 0, @@ -290,17 +255,11 @@ def init_distributed_environment( ), "world group already initialized with a different world size" -_SP: GroupCoordinator | None = None - - def get_sp_group() -> SequenceParallelGroupCoordinator: - assert _SP is not None, "pipeline model parallel group is not initialized" + assert _SP is not None, "sequence parallel group is not initialized" return _SP -_DP: GroupCoordinator | None = None - - def get_dp_group() -> GroupCoordinator: assert _DP is not None, "data parallel group is not initialized" return _DP @@ -472,88 +431,6 @@ def initialize_model_parallel( init_dit_group(dit_parallel_size, backend) -# - - -# def initialize_model_parallel( -# tensor_model_parallel_size: int = 1, -# sequence_model_parallel_size: int = 1, -# data_parallel_size: int = 1, -# backend: str | None = None, -# ) -> None: -# """ -# Initialize model parallel groups. -# -# Arguments: -# tensor_model_parallel_size: number of GPUs used for tensor model -# parallelism (used for language encoder). -# sequence_model_parallel_size: number of GPUs used for sequence model -# parallelism (used for DiT). -# """ -# # Get world size and rank. Ensure some consistencies. -# assert ( -# _WORLD is not None -# ), "world group is not initialized, please call init_distributed_environment first" -# world_size: int = get_world_size() -# backend = backend or torch.distributed.get_backend(get_world_group().device_group) -# assert ( -# world_size >= tensor_model_parallel_size -# ), f"world_size({world_size}) must be greater than or equal to tensor_model_parallel_size({tensor_model_parallel_size})" -# num_tensor_model_parallel_groups: int = world_size // tensor_model_parallel_size -# global _TP -# assert _TP is None, "tensor model parallel group is already initialized" -# group_ranks = [] -# for i in range(num_tensor_model_parallel_groups): -# ranks = list( -# range(i * tensor_model_parallel_size, (i + 1) * tensor_model_parallel_size) -# ) -# group_ranks.append(ranks) -# -# # message queue broadcaster is only used in tensor model parallel group -# _TP = init_parallel_group_coordinator( -# group_ranks, -# get_world_group().local_rank, -# backend, -# use_message_queue_broadcaster=True, -# group_name="tp", -# ) -# -# # Build the sequence model-parallel groups. -# num_sequence_model_parallel_groups: int = world_size // sequence_model_parallel_size -# global _SP -# assert _SP is None, "sequence model parallel group is already initialized" -# group_ranks = [] -# -# # Since SP is incompatible with TP and PP, we can use a simpler group creation logic -# for i in range(num_sequence_model_parallel_groups): -# # Create groups of consecutive ranks -# ranks = list( -# range( -# i * sequence_model_parallel_size, (i + 1) * sequence_model_parallel_size -# ) -# ) -# group_ranks.append(ranks) -# -# _SP = init_parallel_group_coordinator( -# group_ranks, get_world_group().local_rank, backend, group_name="sp" -# ) -# -# # Build the data parallel groups. -# num_data_parallel_groups: int = sequence_model_parallel_size -# global _DP -# assert _DP is None, "data parallel group is already initialized" -# group_ranks = [] -# -# for i in range(num_data_parallel_groups): -# ranks = list(range(i, world_size, num_data_parallel_groups)) -# group_ranks.append(ranks) -# -# _DP = init_parallel_group_coordinator( -# group_ranks, get_world_group().local_rank, backend, group_name="dp" -# ) -# - - def get_sp_world_size() -> int: """Return world size for the sequence model parallel group.""" return get_sp_group().world_size @@ -645,8 +522,14 @@ def maybe_init_distributed_environment_and_model_parallel( def model_parallel_is_initialized() -> bool: - """Check if tensor, sequence parallel groups are initialized.""" - return _TP is not None and _SP is not None and _DP is not None and _CFG is not None + """Check if model parallel groups are initialized.""" + return ( + _DP is not None + and _CFG is not None + and _SP is not None + and _PP is not None + and _TP is not None + ) _TP_STATE_PATCHED = False @@ -795,183 +678,39 @@ def is_the_same_node_as( return [x == 1 for x in aggregated_data.tolist()] -def initialize_tensor_parallel_group( - tensor_model_parallel_size: int = 1, - backend: str | None = None, - group_name_suffix: str = "", -) -> GroupCoordinator: - """Initialize a tensor parallel group for a specific model. - - This function creates a tensor parallel group that can be used with the - patch_tensor_parallel_group context manager. It allows different models - to use different tensor parallelism configurations. - - Arguments: - tensor_model_parallel_size: number of GPUs used for tensor model parallelism. - backend: communication backend to use. - group_name_suffix: optional suffix to make the group name unique. - - Returns: - A GroupCoordinator for tensor parallelism that can be used with - the patch_tensor_parallel_group context manager. - - Example usage: - ```python - # Initialize tensor parallel group for model1 - tp_group_model1 = initialize_tensor_parallel_group( - tensor_model_parallel_size=4, - group_name_suffix="model1" - ) - - # Use tensor parallelism for model1 - with patch_tensor_parallel_group(tp_group_model1): - # Run model1 with tensor parallelism - output1 = model1(input1) - ``` - """ - # Get world size and rank. Ensure some consistencies. - assert torch.distributed.is_initialized() - world_size: int = torch.distributed.get_world_size() - backend = backend or torch.distributed.get_backend(get_world_group().device_group) - - # Ensure the world size is compatible with the parallelism configuration - assert ( - world_size % tensor_model_parallel_size == 0 - ), f"World size ({world_size}) must be divisible by tensor_model_parallel_size ({tensor_model_parallel_size})" - - # Build the tensor model-parallel groups. - num_tensor_model_parallel_groups: int = world_size // tensor_model_parallel_size - tp_group_ranks = [] - for i in range(num_tensor_model_parallel_groups): - ranks = list( - range(i * tensor_model_parallel_size, (i + 1) * tensor_model_parallel_size) - ) - tp_group_ranks.append(ranks) - - # Create TP group coordinator with a unique name - group_name = f"tp_{group_name_suffix}" if group_name_suffix else "tp" - tp_group = init_parallel_group_coordinator( - tp_group_ranks, - get_world_group().local_rank, - backend, - use_message_queue_broadcaster=True, - group_name=group_name, - ) - - return tp_group - - -def initialize_sequence_parallel_group( - sequence_model_parallel_size: int = 1, - backend: str | None = None, - group_name_suffix: str = "", -) -> GroupCoordinator: - """Initialize a sequence parallel group for a specific model. - - This function creates a sequence parallel group that can be used with the - patch_sequence_parallel_group context manager. It allows different models - to use different sequence parallelism configurations. - - Arguments: - sequence_model_parallel_size: number of GPUs used for sequence model parallelism. - backend: communication backend to use. - group_name_suffix: optional suffix to make the group name unique. - - Returns: - A GroupCoordinator for sequence parallelism that can be used with - the patch_sequence_parallel_group context manager. - - Example usage: - ```python - # Initialize sequence parallel group for model2 - sp_group_model2 = initialize_sequence_parallel_group( - sequence_model_parallel_size=2, - group_name_suffix="model2" - ) - - # Use sequence parallelism for model2 - with patch_sequence_parallel_group(sp_group_model2): - # Run model2 with sequence parallelism - output2 = model2(input2) - ``` - """ - # Get world size and rank. Ensure some consistencies. - assert torch.distributed.is_initialized() - world_size: int = torch.distributed.get_world_size() - backend = backend or torch.distributed.get_backend(get_world_group().device_group) - - # Ensure the world size is compatible with the parallelism configuration - assert ( - world_size % sequence_model_parallel_size == 0 - ), f"World size ({world_size}) must be divisible by sequence_model_parallel_size ({sequence_model_parallel_size})" - - # Build the sequence model-parallel groups. - num_sequence_model_parallel_groups: int = world_size // sequence_model_parallel_size - sp_group_ranks = [] - - for i in range(num_sequence_model_parallel_groups): - # Create groups of consecutive ranks - ranks = list( - range( - i * sequence_model_parallel_size, (i + 1) * sequence_model_parallel_size - ) - ) - sp_group_ranks.append(ranks) - - # Create SP group coordinator with a unique name - group_name = f"sp_{group_name_suffix}" if group_name_suffix else "sp" - sp_group = init_parallel_group_coordinator( - sp_group_ranks, get_world_group().local_rank, backend, group_name=group_name - ) - - return sp_group - - -# * QUERY -def get_world_group() -> GroupCoordinator: - assert _WORLD is not None, "world group is not initialized" - return _WORLD - - -# TP -def get_tp_group() -> GroupCoordinator: - assert _TP is not None, "tensor model parallel group is not initialized" - return _TP - - -def get_tensor_model_parallel_world_size(): +def get_tensor_model_parallel_world_size() -> int: """Return world size for the tensor model parallel group.""" - return get_tp_group().world_size + return get_tp_world_size() -def get_tensor_model_parallel_rank(): +def get_tensor_model_parallel_rank() -> int: """Return my rank for the tensor model parallel group.""" - return get_tp_group().rank_in_group + return get_tp_rank() -def get_sequence_parallel_world_size(): +def get_sequence_parallel_world_size() -> int: """Return world size for the sequence parallel group.""" - return get_sp_group().world_size + return get_sp_world_size() -def get_sequence_parallel_rank(): +def get_sequence_parallel_rank() -> int: """Return my rank for the sequence parallel group.""" - return get_sp_group().rank_in_group + return get_sp_parallel_rank() -def get_ulysses_parallel_world_size(): +def get_ulysses_parallel_world_size() -> int: return get_sp_group().ulysses_world_size -def get_ulysses_parallel_rank(): +def get_ulysses_parallel_rank() -> int: return get_sp_group().ulysses_rank -def get_ring_parallel_world_size(): +def get_ring_parallel_world_size() -> int: return get_sp_group().ring_world_size -def get_ring_parallel_rank(): +def get_ring_parallel_rank() -> int: return get_sp_group().ring_rank @@ -981,22 +720,22 @@ def get_pp_group() -> PipelineGroupCoordinator: return _PP -def get_pipeline_parallel_world_size(): +def get_pipeline_parallel_world_size() -> int: """Return world size for the pipeline model parallel group.""" return get_pp_group().world_size -def get_pipeline_parallel_rank(): +def get_pipeline_parallel_rank() -> int: """Return my rank for the pipeline model parallel group.""" return get_pp_group().rank_in_group -def is_pipeline_first_stage(): +def is_pipeline_first_stage() -> bool: """Return True if in the first pipeline model parallel stage, False otherwise.""" return get_pipeline_parallel_rank() == 0 -def is_pipeline_last_stage(): +def is_pipeline_last_stage() -> bool: """Return True if in the last pipeline model parallel stage, False otherwise.""" return get_pipeline_parallel_rank() == (get_pipeline_parallel_world_size() - 1) @@ -1009,33 +748,27 @@ def get_cfg_group() -> GroupCoordinator: return _CFG -def get_classifier_free_guidance_world_size(): +def get_classifier_free_guidance_world_size() -> int: """Return world size for the classifier_free_guidance parallel group.""" return get_cfg_group().world_size -def get_classifier_free_guidance_rank(): +def get_classifier_free_guidance_rank() -> int: """Return my rank for the classifier_free_guidance parallel group.""" return get_cfg_group().rank_in_group -# DP -def get_dp_group() -> GroupCoordinator: - assert _DP is not None, "pipeline model parallel group is not initialized" - return _DP - - -def get_data_parallel_world_size(): +def get_data_parallel_world_size() -> int: """Return world size for the data parallel group.""" - return get_dp_group().world_size + return get_dp_world_size() -def get_data_parallel_rank(): +def get_data_parallel_rank() -> int: """Return my rank for the data parallel group.""" - return get_dp_group().rank_in_group + return get_dp_rank() -def is_dp_last_group(): +def is_dp_last_group() -> bool: """Return True if in the last data parallel group, False otherwise.""" return ( get_sequence_parallel_rank() == (get_sequence_parallel_world_size() - 1) @@ -1045,7 +778,7 @@ def is_dp_last_group(): ) -def get_dit_world_size(): +def get_dit_world_size() -> int: """Return world size for the DiT model (excluding VAE).""" return ( get_data_parallel_world_size() @@ -1056,57 +789,33 @@ def get_dit_world_size(): ) -# Add VAE getter functions -def get_vae_parallel_group() -> GroupCoordinator: +def get_vae_parallel_group() -> ProcessGroup: assert _VAE is not None, "VAE parallel group is not initialized" return _VAE -def get_vae_parallel_world_size(): +def get_vae_parallel_world_size() -> int: """Return world size for the VAE parallel group.""" - return get_vae_parallel_group().world_size + return torch.distributed.get_world_size(group=get_vae_parallel_group()) -def get_vae_parallel_rank(): +def get_vae_parallel_rank() -> int: """Return my rank for the VAE parallel group.""" - return get_vae_parallel_group().rank_in_group - - -# * SET - - -def init_world_group( - ranks: List[int], local_rank: int, backend: str -) -> GroupCoordinator: - return GroupCoordinator( - group_ranks=[ranks], - local_rank=local_rank, - torch_distributed_backend=backend, - ) - - -def model_parallel_is_initialized(): - """Check if tensor and pipeline parallel groups are initialized.""" - return ( - _DP is not None - and _CFG is not None - and _SP is not None - and _PP is not None - and _TP is not None - ) + return torch.distributed.get_rank(group=get_vae_parallel_group()) def init_dit_group( dit_parallel_size: int, backend: str, -): +) -> None: global _DIT + assert _DIT is None, "DIT group is already initialized" _DIT = torch.distributed.new_group( ranks=list(range(dit_parallel_size)), backend=backend ) -def get_dit_group(): +def get_dit_group() -> ProcessGroup: assert _DIT is not None, "DIT group is not initialized" return _DIT @@ -1125,60 +834,14 @@ def init_vae_group( def destroy_model_parallel() -> None: """Set the groups to none and destroy them.""" - global _TP - if _TP: - _TP.destroy() - _TP = None + global _TP, _SP, _DP, _CFG, _PP, _DIT, _VAE - global _SP - if _SP: - _SP.destroy() - _SP = None + for group in (_TP, _SP, _DP, _CFG, _PP): + if group is not None: + group.destroy() - global _DP - if _DP: - _DP.destroy() - _DP = None + for group in (_DIT, _VAE): + if group is not None: + torch.distributed.destroy_process_group(group) - -# xDit -# def destroy_model_parallel(): -# """Set the groups to none and destroy them.""" -# global _DP -# if _DP: -# _DP.destroy() -# _DP = None -# -# global _CFG -# if _CFG: -# _CFG.destroy() -# _CFG = None -# -# global _SP -# if _SP: -# _SP.destroy() -# _SP = None -# -# global _TP -# if _TP: -# _TP.destroy() -# _TP = None -# -# global _PP -# if _PP: -# _PP.destroy() -# _PP = None -# -# global _VAE -# if _VAE: -# _VAE.destroy() -# _VAE = None - - -def destroy_distributed_environment(): - global _WORLD - if _WORLD: - _WORLD.destroy() - _WORLD = None - if torch.distributed.is_initialized(): - torch.distributed.destroy_process_group() + _TP, _SP, _DP, _CFG, _PP, _DIT, _VAE = (None,) * 7