[misc] multiprocess compilation to speed up test (#21483)
This commit is contained in:
@@ -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"]))
|
||||||
|
|||||||
Reference in New Issue
Block a user