Check the topology identities where the layout is written, and build at the published widths (#40340)

This commit is contained in:
Cheng Wan
2026-09-21 12:22:59 -07:00
committed by GitHub
parent d5fdab7022
commit 2d0e94e3a3
43 changed files with 843 additions and 261 deletions
@@ -21,6 +21,7 @@ from sglang.srt.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
)
from sglang.test.test_utils import publish_build_topology
def parse_args():
@@ -85,7 +86,8 @@ def init_dist(backend: str):
distributed_init_method=distributed_init_method,
local_rank=rank,
)
initialize_model_parallel(tensor_model_parallel_size=world_size)
publish_build_topology(world_rank=rank, tp_size=world_size)
initialize_model_parallel()
return dist.group.WORLD
@@ -41,6 +41,7 @@ from sglang.srt.distributed.parallel_state import (
initialize_model_parallel,
set_custom_all_reduce,
)
from sglang.test.test_utils import publish_build_topology
Shape = Tuple[int, int]
@@ -381,7 +382,8 @@ def main():
distributed_init_method="env://",
backend="nccl",
)
initialize_model_parallel(tensor_model_parallel_size=world_size)
publish_build_topology(world_rank=rank, tp_size=world_size)
initialize_model_parallel()
prefill_shapes = parse_shapes(args.prefill_shapes)
decode_shapes = parse_shapes(args.decode_shapes)
@@ -47,6 +47,7 @@ from sglang.srt.distributed.parallel_state import (
initialize_model_parallel,
set_custom_all_reduce,
)
from sglang.test.test_utils import publish_build_topology
Shape = Tuple[int, int]
FP8_DTYPE = torch.float8_e4m3fnuz
@@ -400,7 +401,8 @@ def main() -> None:
distributed_init_method="env://",
backend="nccl",
)
initialize_model_parallel(tensor_model_parallel_size=world_size)
publish_build_topology(world_rank=rank, tp_size=world_size)
initialize_model_parallel()
if rank == 0:
print(
@@ -30,6 +30,7 @@ from sglang.srt.distributed.parallel_state import (
initialize_model_parallel,
set_mscclpp_all_reduce,
)
from sglang.test.test_utils import publish_build_topology
def torch_allreduce(torch_input: torch.Tensor, group: ProcessGroup) -> torch.Tensor:
@@ -173,7 +174,8 @@ if __name__ == "__main__":
rank=rank,
local_rank=rank % 8,
)
initialize_model_parallel(tensor_model_parallel_size=world_size)
publish_build_topology(world_rank=rank, tp_size=world_size)
initialize_model_parallel()
group = get_tensor_model_parallel_group().device_group
cpu_group = get_tensor_model_parallel_group().cpu_group
pynccl_comm = get_tensor_model_parallel_group().pynccl_comm
@@ -44,6 +44,7 @@ from sglang.srt.distributed.parallel_state import (
initialize_model_parallel,
set_torch_symm_mem_all_reduce,
)
from sglang.test.test_utils import publish_build_topology
from sglang.utils import is_in_ci
IS_CI = is_in_ci()
@@ -188,7 +189,8 @@ if __name__ == "__main__":
rank=rank,
local_rank=rank % 8,
)
initialize_model_parallel(tensor_model_parallel_size=world_size)
publish_build_topology(world_rank=rank, tp_size=world_size)
initialize_model_parallel()
group = get_tensor_model_parallel_group().device_group
cpu_group = get_tensor_model_parallel_group().cpu_group
pynccl_comm = get_tensor_model_parallel_group().pynccl_comm