feat(paralell): 添加分布式训练配置与并行工具支持
This commit is contained in:
@@ -0,0 +1,29 @@
|
||||
from khaosz.parallel.utils import (
|
||||
get_world_size,
|
||||
get_rank,
|
||||
get_device_count,
|
||||
get_current_device,
|
||||
get_available_backend,
|
||||
setup_parallel,
|
||||
only_main_procs,
|
||||
spawn_parallel_fn
|
||||
)
|
||||
|
||||
from khaosz.parallel.module import (
|
||||
RowParallelLinear,
|
||||
ColumnParallelLinear
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"get_world_size",
|
||||
"get_rank",
|
||||
"get_device_count",
|
||||
"get_current_device",
|
||||
"get_available_backend",
|
||||
"setup_parallel",
|
||||
"only_main_procs",
|
||||
"spawn_parallel_fn",
|
||||
|
||||
"RowParallelLinear",
|
||||
"ColumnParallelLinear"
|
||||
]
|
||||
@@ -33,6 +33,17 @@ def get_available_backend():
|
||||
else:
|
||||
return "gloo"
|
||||
|
||||
def get_world_size() -> int:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
return dist.get_world_size()
|
||||
else:
|
||||
return 1
|
||||
|
||||
def get_rank() -> int:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
return dist.get_rank()
|
||||
else:
|
||||
return 0
|
||||
|
||||
@contextmanager
|
||||
def setup_parallel(
|
||||
@@ -76,6 +87,21 @@ def setup_parallel(
|
||||
if dist.is_initialized():
|
||||
dist.destroy_process_group()
|
||||
|
||||
@contextmanager
|
||||
def only_main_procs(main_process_rank=0, block=True):
|
||||
is_main_proc = (get_rank() == main_process_rank)
|
||||
|
||||
if dist.is_initialized() and block:
|
||||
dist.barrier()
|
||||
|
||||
try:
|
||||
yield is_main_proc
|
||||
|
||||
finally:
|
||||
if dist.is_initialized() and block:
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def wrapper_spawn_func(rank, world_size, func, kwargs_dict):
|
||||
with setup_parallel(rank, world_size):
|
||||
func(**kwargs_dict)
|
||||
|
||||
Reference in New Issue
Block a user