[misc] multiprocess compilation to speed up test (#21483)

This commit is contained in:
DarkSharpness
2026-03-31 08:56:37 +08:00
committed by GitHub
parent 3650bfb199
commit 4e480982fa
2 changed files with 41 additions and 19 deletions
+2 -2
View File
@@ -95,7 +95,7 @@ if TYPE_CHECKING:
def _jit_custom_all_reduce_pull_module(dtype: torch.dtype, world_size: int): def _jit_custom_all_reduce_pull_module(dtype: torch.dtype, world_size: int):
args = make_cpp_args(dtype, world_size, is_arch_support_pdl()) args = make_cpp_args(dtype, world_size, is_arch_support_pdl())
return load_jit( return load_jit(
"custom_all_reduce", "custom_all_reduce_pull",
*args, *args,
extra_ldflags=["-lcuda"], extra_ldflags=["-lcuda"],
cuda_files=["distributed/custom_all_reduce_pull.cuh"], cuda_files=["distributed/custom_all_reduce_pull.cuh"],
@@ -107,7 +107,7 @@ def _jit_custom_all_reduce_pull_module(dtype: torch.dtype, world_size: int):
def _jit_custom_all_reduce_push_module(dtype: torch.dtype, world_size: int): def _jit_custom_all_reduce_push_module(dtype: torch.dtype, world_size: int):
args = make_cpp_args(dtype, world_size, is_arch_support_pdl()) args = make_cpp_args(dtype, world_size, is_arch_support_pdl())
return load_jit( return load_jit(
"custom_all_reduce", "custom_all_reduce_push",
*args, *args,
extra_ldflags=["-lcuda"], extra_ldflags=["-lcuda"],
cuda_files=["distributed/custom_all_reduce_push.cuh"], cuda_files=["distributed/custom_all_reduce_push.cuh"],
@@ -16,29 +16,33 @@ from __future__ import annotations
import itertools import itertools
import logging import logging
import multiprocessing as mp
import os import os
import subprocess import subprocess
import sys import sys
from typing import Optional from typing import Dict, Optional, Tuple
import pytest import pytest
import torch import torch
import torch.distributed as dist import torch.distributed as dist
from tqdm import tqdm
import sglang.srt.distributed.parallel_state as ps import sglang.srt.distributed.parallel_state as ps
from sglang.jit_kernel.all_reduce import AllReduceAlgo from sglang.jit_kernel.all_reduce import (
AllReduceAlgo,
_jit_custom_all_reduce_pull_module,
_jit_custom_all_reduce_push_module,
)
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2, CustomAllReduceV2,
) )
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci( register_cuda_ci(
est_time=500, est_time=300,
suite="stage-b-kernel-unit-8-gpu-h200", suite="stage-b-kernel-unit-8-gpu-h200",
) )
register_cuda_ci( register_cuda_ci(
est_time=500, est_time=300,
suite="nightly-kernel-8-gpu-h200", suite="nightly-kernel-8-gpu-h200",
nightly=True, nightly=True,
) )
@@ -67,7 +71,7 @@ SHOTS = [
] ]
USE_GRAPH_OPTIONS = [True, False] USE_GRAPH_OPTIONS = [True, False]
TEST_CONFIG = itertools.product(TEST_SIZES, TEST_DTYPES, SHOTS, USE_GRAPH_OPTIONS) TEST_CONFIG = itertools.product(TEST_SIZES, TEST_DTYPES, SHOTS, USE_GRAPH_OPTIONS)
TEST_LAYERS = 2 TEST_LAYERS = 4
TEST_LOOP = 16 TEST_LOOP = 16
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -75,14 +79,13 @@ TEST_LOOP = 16
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def run_torchrun(nproc: int, timeout: int = 300) -> None: def _run_torchrun(nproc: int, timeout: int = 300) -> None:
"""Launch this script as a torchrun worker and assert success.""" """Launch this script as a torchrun worker and assert success."""
cmd = [ cmd = [
"torchrun", "torchrun",
f"--nproc_per_node={nproc}", f"--nproc_per_node={nproc}",
__file__, __file__,
] ]
os.environ["DISABLE_PBAR"] = "1"
result = subprocess.run( result = subprocess.run(
cmd, cmd,
stdout=subprocess.PIPE, stdout=subprocess.PIPE,
@@ -96,14 +99,37 @@ def run_torchrun(nproc: int, timeout: int = 300) -> None:
) )
@pytest.mark.parametrize("nproc", [2, 3, 4, 5, 6, 7, 8]) def _compile_one(dtype: torch.dtype, world_size: int):
_jit_custom_all_reduce_push_module(dtype, world_size)
_jit_custom_all_reduce_pull_module(dtype, world_size)
def _precompile_kernels() -> None:
# NOTE: even when device count < 8, we should be able to compile all
process_map: Dict[Tuple[torch.dtype, int], mp.Process] = {}
COMPILE_SPACE = itertools.product(TEST_DTYPES, [2, 3, 4, 5, 6, 7, 8])
mp.set_start_method("spawn")
for config in COMPILE_SPACE:
process_map[config] = mp.Process(target=_compile_one, args=config)
for process in process_map.values():
process.start()
for (dtype, world_size), process in process_map.items():
process.join()
if process.exitcode != 0:
raise RuntimeError(f"Custom All Reduce {world_size=} {dtype=} failed")
@pytest.mark.parametrize("nproc", [1, 2, 3, 4, 5, 6, 7, 8])
def test_custom_allreduce(nproc: int) -> None: def test_custom_allreduce(nproc: int) -> None:
if nproc == 1: # NOTE: special case to speed up tests
return _precompile_kernels()
device_count = torch.cuda.device_count() device_count = torch.cuda.device_count()
if device_count < nproc: if device_count < nproc:
pytest.skip( pytest.skip(
f"Requires at least {nproc} GPUs, but only {device_count} available" f"Requires at least {nproc} GPUs, but only {device_count} available"
) )
run_torchrun(nproc) _run_torchrun(nproc)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -192,8 +218,6 @@ def worker_test(
dist.all_reduce(out_ref, group=nccl_group) dist.all_reduce(out_ref, group=nccl_group)
out_jit = run_fn(inp) out_jit = run_fn(inp)
num_errors += not torch.all(out_jit == out_ref) num_errors += not torch.all(out_jit == out_ref)
torch.cuda.synchronize()
nccl_group.barrier().wait()
if num_errors > 0: if num_errors > 0:
return RuntimeError( return RuntimeError(
f"Test failed for {size=}, {dtype=}, {algo=}, " f"Test failed for {size=}, {dtype=}, {algo=}, "
@@ -211,9 +235,7 @@ def worker_main() -> None:
logging.disable(logging.INFO) # Suppress internal logging for cleaner test output logging.disable(logging.INFO) # Suppress internal logging for cleaner test output
items = list(enumerate(TEST_CONFIG)) items = list(enumerate(TEST_CONFIG))
disable_pbar = os.environ.get("DISABLE_PBAR", "0") == "1" or rank != 0 for i, (size, dtype, algo, use_graph) in items:
pbar = tqdm(items, desc=f"Testing {world_size} GPUs", disable=disable_pbar)
for i, (size, dtype, algo, use_graph) in pbar:
error = worker_test(device, nccl_group, comm, size, dtype, use_graph, algo) error = worker_test(device, nccl_group, comm, size, dtype, use_graph, algo)
if error is not None: if error is not None:
print( print(
@@ -222,7 +244,7 @@ def worker_main() -> None:
f"Error: {error}" f"Error: {error}"
) )
# communicate the result to rank 0 for logging # communicate the result to rank 0 for logging
result = torch.tensor([int(error is not None)], device=device) result = torch.tensor([int(error is not None)])
dist.all_reduce(result, group=cpu_group) dist.all_reduce(result, group=cpu_group)
failed = bool(result.item()) failed = bool(result.item())
if failed: if failed:
@@ -239,4 +261,4 @@ if __name__ == "__main__":
if "LOCAL_RANK" in os.environ: if "LOCAL_RANK" in os.environ:
worker_main() worker_main()
else: else:
sys.exit(pytest.main([__file__, "-v", "-s"])) sys.exit(pytest.main([__file__, "-x", "-vv", "-s"]))