[diffusion] refactor: cleanup parallel_state.py (#20760)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot]
parent
17c81a3e07
commit
5717834f1f
@@ -59,14 +59,14 @@ from .group_coordinator import (
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
_WORLD: Optional[GroupCoordinator] = None
|
_WORLD: GroupCoordinator | None = None
|
||||||
_TP: Optional[GroupCoordinator] = None
|
_TP: GroupCoordinator | None = None
|
||||||
_SP: Optional[SequenceParallelGroupCoordinator] = None
|
_SP: SequenceParallelGroupCoordinator | None = None
|
||||||
_PP: Optional[PipelineGroupCoordinator] = None
|
_PP: PipelineGroupCoordinator | None = None
|
||||||
_CFG: Optional[GroupCoordinator] = None
|
_CFG: GroupCoordinator | None = None
|
||||||
_DP: Optional[GroupCoordinator] = None
|
_DP: GroupCoordinator | None = None
|
||||||
_DIT: Optional[GroupCoordinator] = None
|
_DIT: ProcessGroup | None = None
|
||||||
_VAE: Optional[GroupCoordinator] = None
|
_VAE: ProcessGroup | None = None
|
||||||
|
|
||||||
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
|
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)
|
return torch.empty_like(tensor)
|
||||||
|
|
||||||
|
|
||||||
_WORLD: GroupCoordinator | None = None
|
|
||||||
_NODE: GroupCoordinator | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def get_world_group() -> GroupCoordinator:
|
def get_world_group() -> GroupCoordinator:
|
||||||
assert _WORLD is not None, "world group is not initialized"
|
assert _WORLD is not None, "world group is not initialized"
|
||||||
return _WORLD
|
return _WORLD
|
||||||
@@ -137,7 +133,6 @@ def init_world_group(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# xDiT
|
|
||||||
def init_parallel_group_coordinator(
|
def init_parallel_group_coordinator(
|
||||||
group_ranks: List[List[int]],
|
group_ranks: List[List[int]],
|
||||||
local_rank: int,
|
local_rank: int,
|
||||||
@@ -145,9 +140,7 @@ def init_parallel_group_coordinator(
|
|||||||
parallel_mode: str,
|
parallel_mode: str,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> GroupCoordinator:
|
) -> GroupCoordinator:
|
||||||
"""
|
"""Return a group coordinator for the given parallel mode."""
|
||||||
Returns a Group Coordinator for the given parallel mode
|
|
||||||
"""
|
|
||||||
assert parallel_mode in [
|
assert parallel_mode in [
|
||||||
"data",
|
"data",
|
||||||
"pipeline",
|
"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:
|
def get_tp_group() -> GroupCoordinator:
|
||||||
assert _TP is not None, "tensor model parallel group is not initialized"
|
assert _TP is not None, "tensor model parallel group is not initialized"
|
||||||
return _TP
|
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(
|
def init_distributed_environment(
|
||||||
world_size: int = 1,
|
world_size: int = 1,
|
||||||
rank: int = 0,
|
rank: int = 0,
|
||||||
@@ -290,17 +255,11 @@ def init_distributed_environment(
|
|||||||
), "world group already initialized with a different world size"
|
), "world group already initialized with a different world size"
|
||||||
|
|
||||||
|
|
||||||
_SP: GroupCoordinator | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def get_sp_group() -> SequenceParallelGroupCoordinator:
|
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
|
return _SP
|
||||||
|
|
||||||
|
|
||||||
_DP: GroupCoordinator | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def get_dp_group() -> GroupCoordinator:
|
def get_dp_group() -> GroupCoordinator:
|
||||||
assert _DP is not None, "data parallel group is not initialized"
|
assert _DP is not None, "data parallel group is not initialized"
|
||||||
return _DP
|
return _DP
|
||||||
@@ -472,88 +431,6 @@ def initialize_model_parallel(
|
|||||||
init_dit_group(dit_parallel_size, backend)
|
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:
|
def get_sp_world_size() -> int:
|
||||||
"""Return world size for the sequence model parallel group."""
|
"""Return world size for the sequence model parallel group."""
|
||||||
return get_sp_group().world_size
|
return get_sp_group().world_size
|
||||||
@@ -645,8 +522,14 @@ def maybe_init_distributed_environment_and_model_parallel(
|
|||||||
|
|
||||||
|
|
||||||
def model_parallel_is_initialized() -> bool:
|
def model_parallel_is_initialized() -> bool:
|
||||||
"""Check if tensor, sequence parallel groups are initialized."""
|
"""Check if model parallel groups are initialized."""
|
||||||
return _TP is not None and _SP is not None and _DP is not None and _CFG is not None
|
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
|
_TP_STATE_PATCHED = False
|
||||||
@@ -795,183 +678,39 @@ def is_the_same_node_as(
|
|||||||
return [x == 1 for x in aggregated_data.tolist()]
|
return [x == 1 for x in aggregated_data.tolist()]
|
||||||
|
|
||||||
|
|
||||||
def initialize_tensor_parallel_group(
|
def get_tensor_model_parallel_world_size() -> int:
|
||||||
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():
|
|
||||||
"""Return world size for the tensor model parallel group."""
|
"""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 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 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 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
|
return get_sp_group().ulysses_world_size
|
||||||
|
|
||||||
|
|
||||||
def get_ulysses_parallel_rank():
|
def get_ulysses_parallel_rank() -> int:
|
||||||
return get_sp_group().ulysses_rank
|
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
|
return get_sp_group().ring_world_size
|
||||||
|
|
||||||
|
|
||||||
def get_ring_parallel_rank():
|
def get_ring_parallel_rank() -> int:
|
||||||
return get_sp_group().ring_rank
|
return get_sp_group().ring_rank
|
||||||
|
|
||||||
|
|
||||||
@@ -981,22 +720,22 @@ def get_pp_group() -> PipelineGroupCoordinator:
|
|||||||
return _PP
|
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 world size for the pipeline model parallel group."""
|
||||||
return get_pp_group().world_size
|
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 my rank for the pipeline model parallel group."""
|
||||||
return get_pp_group().rank_in_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 True if in the first pipeline model parallel stage, False otherwise."""
|
||||||
return get_pipeline_parallel_rank() == 0
|
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 True if in the last pipeline model parallel stage, False otherwise."""
|
||||||
return get_pipeline_parallel_rank() == (get_pipeline_parallel_world_size() - 1)
|
return get_pipeline_parallel_rank() == (get_pipeline_parallel_world_size() - 1)
|
||||||
|
|
||||||
@@ -1009,33 +748,27 @@ def get_cfg_group() -> GroupCoordinator:
|
|||||||
return _CFG
|
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 world size for the classifier_free_guidance parallel group."""
|
||||||
return get_cfg_group().world_size
|
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 my rank for the classifier_free_guidance parallel group."""
|
||||||
return get_cfg_group().rank_in_group
|
return get_cfg_group().rank_in_group
|
||||||
|
|
||||||
|
|
||||||
# DP
|
def get_data_parallel_world_size() -> int:
|
||||||
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():
|
|
||||||
"""Return world size for the data parallel group."""
|
"""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 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 True if in the last data parallel group, False otherwise."""
|
||||||
return (
|
return (
|
||||||
get_sequence_parallel_rank() == (get_sequence_parallel_world_size() - 1)
|
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 world size for the DiT model (excluding VAE)."""
|
||||||
return (
|
return (
|
||||||
get_data_parallel_world_size()
|
get_data_parallel_world_size()
|
||||||
@@ -1056,57 +789,33 @@ def get_dit_world_size():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Add VAE getter functions
|
def get_vae_parallel_group() -> ProcessGroup:
|
||||||
def get_vae_parallel_group() -> GroupCoordinator:
|
|
||||||
assert _VAE is not None, "VAE parallel group is not initialized"
|
assert _VAE is not None, "VAE parallel group is not initialized"
|
||||||
return _VAE
|
return _VAE
|
||||||
|
|
||||||
|
|
||||||
def get_vae_parallel_world_size():
|
def get_vae_parallel_world_size() -> int:
|
||||||
"""Return world size for the VAE parallel group."""
|
"""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 my rank for the VAE parallel group."""
|
||||||
return get_vae_parallel_group().rank_in_group
|
return torch.distributed.get_rank(group=get_vae_parallel_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
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def init_dit_group(
|
def init_dit_group(
|
||||||
dit_parallel_size: int,
|
dit_parallel_size: int,
|
||||||
backend: str,
|
backend: str,
|
||||||
):
|
) -> None:
|
||||||
global _DIT
|
global _DIT
|
||||||
|
assert _DIT is None, "DIT group is already initialized"
|
||||||
_DIT = torch.distributed.new_group(
|
_DIT = torch.distributed.new_group(
|
||||||
ranks=list(range(dit_parallel_size)), backend=backend
|
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"
|
assert _DIT is not None, "DIT group is not initialized"
|
||||||
return _DIT
|
return _DIT
|
||||||
|
|
||||||
@@ -1125,60 +834,14 @@ def init_vae_group(
|
|||||||
|
|
||||||
def destroy_model_parallel() -> None:
|
def destroy_model_parallel() -> None:
|
||||||
"""Set the groups to none and destroy them."""
|
"""Set the groups to none and destroy them."""
|
||||||
global _TP
|
global _TP, _SP, _DP, _CFG, _PP, _DIT, _VAE
|
||||||
if _TP:
|
|
||||||
_TP.destroy()
|
|
||||||
_TP = None
|
|
||||||
|
|
||||||
global _SP
|
for group in (_TP, _SP, _DP, _CFG, _PP):
|
||||||
if _SP:
|
if group is not None:
|
||||||
_SP.destroy()
|
group.destroy()
|
||||||
_SP = None
|
|
||||||
|
|
||||||
global _DP
|
for group in (_DIT, _VAE):
|
||||||
if _DP:
|
if group is not None:
|
||||||
_DP.destroy()
|
torch.distributed.destroy_process_group(group)
|
||||||
_DP = None
|
|
||||||
|
|
||||||
|
_TP, _SP, _DP, _CFG, _PP, _DIT, _VAE = (None,) * 7
|
||||||
# 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()
|
|
||||||
|
|||||||
Reference in New Issue
Block a user