Diffusers
Safetensors
HY / trainer /distributed /__init__.py
Cccccz's picture
Upload batch 65: 500 files (0.01 GiB)
74da989 verified
Raw History Blame Contribute Delete
1.27 kB
# SPDX-License-Identifier: Apache-2.0
from trainer.distributed.communication_op import *
from trainer.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
get_local_torch_device, get_sp_group, get_sp_parallel_rank,
get_sp_world_size, get_tp_group, get_tp_rank, get_tp_world_size,
get_world_group, get_world_rank, get_world_size,
init_distributed_environment, initialize_model_parallel,
maybe_init_distributed_environment_and_model_parallel,
model_parallel_is_initialized)
from trainer.distributed.utils import *
__all__ = [
# Initialization
"init_distributed_environment",
"initialize_model_parallel",
"cleanup_dist_env_and_memory",
"model_parallel_is_initialized",
"maybe_init_distributed_environment_and_model_parallel",
# World group
"get_world_group",
"get_world_rank",
"get_world_size",
# Data parallel group
"get_dp_group",
"get_dp_rank",
"get_dp_world_size",
# Sequence parallel group
"get_sp_group",
"get_sp_parallel_rank",
"get_sp_world_size",
# Tensor parallel group
"get_tp_group",
"get_tp_rank",
"get_tp_world_size",
# Get torch device
"get_local_torch_device",
]