[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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user