feat(paralell): 添加分布式训练配置与并行工具支持

This commit is contained in:
2025-12-05 13:52:17 +08:00
parent d31137a2db
commit d52685facd
4 changed files with 72 additions and 1 deletions
+29
View File
@@ -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"
]
+26
View File
@@ -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)