feat: 增加server, 并且修改测试单元

This commit is contained in:
2026-04-02 15:05:07 +08:00
parent 9f1561afe7
commit 475de51c7d
12 changed files with 616 additions and 99 deletions
+32
View File
@@ -0,0 +1,32 @@
import torch
import torch.distributed as dist
from astrai.parallel import get_rank, only_on_rank, spawn_parallel_fn
@only_on_rank(0)
def _test_only_on_rank_helper():
return True
def only_on_rank():
result = _test_only_on_rank_helper()
if get_rank() == 0:
assert result is True
else:
assert result is None
def all_reduce():
x = torch.tensor([get_rank()], dtype=torch.int)
dist.all_reduce(x, op=dist.ReduceOp.SUM)
expected_sum = sum(range(dist.get_world_size()))
assert x.item() == expected_sum
def test_spawn_only_on_rank():
spawn_parallel_fn(only_on_rank, world_size=2, backend="gloo")
def test_spawn_all_reduce():
spawn_parallel_fn(all_reduce, world_size=2, backend="gloo")