[JIT Kernel] Multi-GPU test/bench framework for custom all-reduce + TP QKNorm (#26706)

Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: ziyi.xu <ziyi.xu@radixark.ai>
This commit is contained in:
DarkSharpness
2026-06-14 17:20:36 +08:00
committed by GitHub
co-authored by Claude ziyi.xu
parent 5331de0f8c
commit d72314808f
8 changed files with 1044 additions and 661 deletions
+144 -37
View File
@@ -1,3 +1,4 @@
import contextlib
import inspect
import itertools
import math
@@ -5,6 +6,7 @@ import os
from typing import (
Any,
Callable,
ContextManager,
Dict,
Generic,
Iterable,
@@ -93,6 +95,60 @@ def _process_metrics(times: list[float], metrics: tuple[Metric, ...]) -> list[fl
return results
@cache_once
def _get_l2_cache_size() -> int:
device = torch.cuda.current_device()
props = torch.cuda.get_device_properties(device)
return props.L2_cache_size
_L2_SAFE_RATIO = 5
def _get_flush_l2_buffer() -> torch.Tensor:
"""Get a buffer sized to flush the L2 cache when accessed."""
device = torch.device(f"cuda:{torch.cuda.current_device()}")
l2_size = _get_l2_cache_size()
safe_size = int(l2_size * _L2_SAFE_RATIO)
return torch.empty(safe_size, device=device, dtype=torch.uint8)
def _calculate_rotation_count(nbytes: int, min_rotations: int = 2) -> int:
"""
Adapted from flashinfer benchmark utility:
https://github.com/flashinfer-ai/flashinfer/blob/c5a2b06edae4fa2bfd2ae25eed16eb565c70513f/flashinfer/testing/utils.py
Calculate the number of buffer copies needed to ensure cold L2 cache.
The function uses conservative thresholds to account for:
- LRU eviction being gradual (not all data evicted when capacity exceeded)
- Cache associativity effects (some data may persist in non-conflicting sets)
- Hardware prefetching behavior
Returns 1 (no rotation needed) only when tensor size substantially exceeds
L2 cache, ensuring cache effects are truly negligible.
Args:
tensors: List of tensors to consider for rotation (must be on GPU).
device: Device for L2 cache query (None for current device).
min_rotations: Minimum number of rotations when rotation is needed.
Returns:
Number of buffer copies needed (1 means no rotation needed).
"""
l2_size = _get_l2_cache_size()
safe_cache_threshold = l2_size * _L2_SAFE_RATIO
if nbytes <= 0 or nbytes >= safe_cache_threshold:
return 1 # No tensors to rotate
# Conservative formula: ensure between any two uses of the same buffer,
# we've accessed enough data to fully flush L2 with margin
# Using safe_cache_threshold ensures we account for all cache effects
num_rotations = math.ceil(safe_cache_threshold / nbytes) + 1
return max(min_rotations, num_rotations)
class BenchResult(NamedTuple):
metrics: Tuple[Metric, ...]
times: List[float] # in seconds
@@ -319,6 +375,67 @@ def parametrize(names: str, vals: List[Any], ci_vals: Optional[List[Any]] = None
return decorator
def _do_bench_internal_graph(
fn: Callable,
replay_iters: int,
input_args: Tuple[Any, ...],
input_kwargs: Dict[str, Any],
graph_clone_args: Iterable[int],
graph_clone_kwargs: Iterable[str],
graph_context: ContextManager,
sync_multigpu_fn: Callable[[], Any],
) -> List[float]:
result: List[float] = []
stream = torch.cuda.current_stream()
empty_tensor = _get_flush_l2_buffer()
# only count the cloned tensors for rotation count
nbytes = sum(_get_nbytes_recursive(input_args[i]) for i in graph_clone_args)
nbytes += sum(_get_nbytes_recursive(input_kwargs[k]) for k in graph_clone_kwargs)
rotate_count = min(_calculate_rotation_count(nbytes), 100)
loop_count = math.ceil(100 / rotate_count) * rotate_count
input_args_list = [input_args] * rotate_count
input_kwargs_list = [input_kwargs] * rotate_count
graph_clone_args = set(graph_clone_args)
graph_clone_kwargs = set(graph_clone_kwargs)
graph = torch.cuda.CUDAGraph()
# NOTE: we rotate the buffer here to avoid L2 cache effect
for i in range(1, rotate_count):
input_args_list[i] = tuple(
(
_clone_recursive(input_args[j])
if j in graph_clone_args
else input_args[j]
)
for j in range(len(input_args))
)
input_kwargs_list[i] = dict(
(k, (_clone_recursive(v) if k in graph_clone_kwargs else v))
for k, v in input_kwargs.items()
)
with graph_context:
with torch.cuda.graph(graph, stream=stream):
for i in range(loop_count):
args = input_args_list[i % rotate_count]
kwargs = input_kwargs_list[i % rotate_count]
fn(*args, **kwargs)
# warm up the graph once
graph.replay()
# then replay the graph and measure the time
tic = torch.cuda.Event(enable_timing=True)
toc = torch.cuda.Event(enable_timing=True)
for _ in range(max(replay_iters // loop_count, 10)):
empty_tensor.zero_() # cold the L2 cache
sync_multigpu_fn() # sync GPU before each iteration for precise timing
tic.record(stream)
graph.replay()
toc.record(stream)
stream.synchronize()
result.append(tic.elapsed_time(toc) / loop_count)
return result
def do_bench(
fn: Callable,
*,
@@ -338,9 +455,13 @@ def do_bench(
memory_output: Iterable[Any] | Literal["out"] | None = "out",
extra_memory_args: Iterable[Any] | None = None,
extra_memory_footprint: int = 0,
graph_context_fn: Optional[Callable[[], ContextManager]] = None,
sync_multigpu_fn: Optional[Callable[[], Any]] = None,
) -> BenchResult:
"""
Benchmark a function using CUDA graph or naive loop.
Adapted from flashinfer benchmark utility:
https://github.com/flashinfer-ai/flashinfer/blob/c5a2b06edae4fa2bfd2ae25eed16eb565c70513f/flashinfer/testing/utils.py
:param fn: Function to benchmark
:param input_args: Positional arguments to pass to the function
@@ -361,6 +482,10 @@ def do_bench(
:param extra_memory_args: Additional arguments to consider for memory footprint calculation.
:param extra_memory_footprint: Additional memory footprint to consider.
This is typically used when the load/store bytes is dynamic.
:param graph_context_fn: A callable returning a context manager that wraps the cuda graph capture.
:param sync_multigpu_fn: A callable to synchronize multiple GPUs before each iteration. For precise
benchmark number in multi-GPU benchmark, it should be some synchronization
primitive on GPU side (not on CPU side).
"""
# first warmup the function
device_id = torch.cuda.current_device()
@@ -368,17 +493,14 @@ def do_bench(
stream = _get_benchmark_stream(device_id)
old_current_stream = torch.cuda.current_stream(device_id)
result: List[float] = []
sync_multigpu_fn = sync_multigpu_fn or (lambda: None)
with torch.cuda.device(device_id), torch.cuda.stream(stream):
stream.wait_stream(old_current_stream)
sync_multigpu_fn()
for _ in range(warmup_iters):
fn(*input_args, **input_kwargs)
if use_cuda_graph:
# NOTE: by default, reduce all the CPU-side overhead
rep_count = 4
loop_iters = 100
graph = torch.cuda.CUDAGraph()
input_args_list = [input_args] * rep_count
input_kwargs_list = [input_kwargs] * rep_count
if graph_clone_args == "all":
graph_clone_args = range(len(input_args))
elif graph_clone_args is None:
@@ -387,44 +509,29 @@ def do_bench(
graph_clone_kwargs = input_kwargs.keys()
elif graph_clone_kwargs is None:
graph_clone_kwargs = []
graph_clone_args = set(graph_clone_args)
graph_clone_kwargs = set(graph_clone_kwargs)
# NOTE: we rotate the buffer here to avoid L2 cache effect
for i in range(1, rep_count):
input_args_list[i] = tuple(
(
_clone_recursive(input_args[j])
if j in graph_clone_args
else input_args[j]
)
for j in range(len(input_args))
)
input_kwargs_list[i] = dict(
(k, (_clone_recursive(v) if k in graph_clone_kwargs else v))
for k, v in input_kwargs.items()
)
with torch.cuda.graph(graph, stream=stream):
for _ in range(loop_iters // rep_count):
for args, kwargs in zip(input_args_list, input_kwargs_list):
fn(*args, **kwargs)
# warm up the graph
graph.replay()
# then replay the graph and measure the time
tic = torch.cuda.Event(enable_timing=True)
toc = torch.cuda.Event(enable_timing=True)
for _ in range(max(replay_iters // loop_iters, 10)):
tic.record(stream)
graph.replay()
toc.record(stream)
stream.synchronize()
result.append(tic.elapsed_time(toc) / loop_iters)
graph_context = (
graph_context_fn()
if graph_context_fn is not None
else contextlib.nullcontext()
)
result = _do_bench_internal_graph(
fn,
replay_iters,
input_args,
input_kwargs,
graph_clone_args,
graph_clone_kwargs,
graph_context,
sync_multigpu_fn,
)
else:
# NOTE: no cuda graph, naive loop
empty_tensor = torch.empty(64 * 1024 * 1024, device=f"cuda:{device_id}")
tic = torch.cuda.Event(enable_timing=True)
toc = torch.cuda.Event(enable_timing=True)
empty_tensor = _get_flush_l2_buffer()
for _ in range(max(replay_iters, 10)):
empty_tensor.zero_() # cold the L2 cache
sync_multigpu_fn()
tic.record(stream)
fn(*input_args, **input_kwargs)
toc.record(stream)
+53 -1
View File
@@ -1,12 +1,64 @@
"""Common utilities for jit_kernel benchmark files."""
from typing import Callable, List, Sequence, Tuple
from typing import Callable, List, Optional, Sequence, Tuple
import torch
import triton.testing
from sglang.jit_kernel.mp import multigpu_launch
from sglang.utils import is_in_ci
def multigpu_bench_main(
name: str,
file: str,
num_gpus: Sequence[int],
main_fn: Callable[[], None],
*,
pre_launch_fn: Optional[Callable[[List[int]], None]] = None,
timeout: Optional[int] = None,
) -> None:
"""cudalib-style multi-GPU benchmark entry point.
Drop this at the bottom of a benchmark file::
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=range(2, 9),
main_fn=benchmark.run,
)
Mirrors :func:`multigpu_pytest_main` but invokes a caller-supplied function
instead of pytest. ``main_fn`` is expected to return ``None`` on success;
any exception propagates as a non-zero exit. Pass ``--num-gpu 2,4`` on the
command line to override ``num_gpus``.
``pre_launch_fn`` (kw-only) runs once in the outer process before any
torchrun child starts, receiving the runnable world sizes. Use it for
parallel JIT precompilation so torchrun children hit a warm disk cache.
``timeout`` (kw-only, seconds) bounds each per-world-size torchrun
invocation. Defaults to ``None`` (wait indefinitely) since benchmark sweeps
can legitimately run long; set it to fail fast on a hung worker.
"""
def inner() -> int:
main_fn()
return 0
return multigpu_launch(
name,
file,
num_gpus,
env_key="_IS_BENCH_MULTIGPU_SGLANG_JIT_KERNEL",
inner=inner,
kind="benchmark",
pre_launch_fn=pre_launch_fn,
timeout=timeout,
)
# Common constants
DEFAULT_DTYPE = torch.bfloat16
DEFAULT_DEVICE = "cuda"
+214
View File
@@ -0,0 +1,214 @@
"""Multi-process / multi-GPU launching utilities (torchrun-based).
Shared `multigpu_launch` helper that both `sglang.jit_kernel.tests.utils` and
`sglang.jit_kernel.benchmark.utils` build their domain-specific entry points on
top of (`multigpu_pytest_main`, `multigpu_bench_main`).
When a script that calls one of those wrappers is run with plain `python`, the
launcher relaunches the same file under `torchrun` once for each `N` in
`num_gpus`. When the inner workers run (identified by an env_key being set),
the same launcher calls `inner()` on every rank, silences stdout on non-zero
ranks, and exits with its return code.
"""
from __future__ import annotations
import atexit
import logging
import os
import signal
import subprocess
import sys
from typing import Any, Callable, List, NoReturn, Optional, Sequence
import psutil
import torch
logger = logging.getLogger(__name__)
def register_comm_cleanup(comm: Any) -> None:
"""Register an idempotent shutdown for a custom-AR communicator."""
def _safe_close() -> None:
try:
comm.close()
except Exception:
pass
# Disable both class flavors' early-out paths in __del__/close.
try:
comm.disabled = True
except Exception:
pass
# CustomAllReduceV2: drop ``obj`` so close() short-circuits next time.
try:
delattr(comm, "obj")
except Exception:
pass
# CustomAllreduce: zero ``_ptr`` so close() short-circuits next time.
try:
comm._ptr = 0
except Exception:
pass
atexit.register(_safe_close)
def _kill_pgroup(pgid: int) -> None:
try:
os.killpg(pgid, signal.SIGKILL)
except ProcessLookupError:
pass
def _kill_descendants(pid: int) -> None:
"""Snapshot every descendant of `pid` *now* and SIGKILL them all.
Must be called BEFORE the direct child (torchrun) dies -- once it does,
its workers get reparented to init and we lose them via the process tree.
"""
try:
root = psutil.Process(pid)
except psutil.NoSuchProcess:
return
descendants = root.children(recursive=True)
for proc in descendants:
try:
proc.kill()
except psutil.Error:
# NoSuchProcess (already gone) or AccessDenied -- nothing to do.
pass
psutil.wait_procs(descendants, timeout=5)
def _extract_num_gpus_override(
argv: list[str],
) -> tuple[list[int] | None, list[str]]:
"""Pop `--num-gpu(s)` flags out of `argv` and return them separately.
Accepts `--num-gpu N`, `--num-gpu=N`, `--num-gpus ...`, and comma-separated
lists like `--num-gpu 2,4,8`. May be repeated.
"""
override: list[int] = []
remaining: list[str] = []
i = 0
while i < len(argv):
a = argv[i]
if a in ("--num-gpu", "--num-gpus"):
if i + 1 >= len(argv):
raise ValueError(f"missing value for {a} (expected e.g. `{a} 2,4`)")
override.extend(int(x) for x in argv[i + 1].split(","))
i += 2
elif a.startswith("--num-gpu=") or a.startswith("--num-gpus="):
_, val = a.split("=", 1)
override.extend(int(x) for x in val.split(","))
i += 1
else:
remaining.append(a)
i += 1
return (override if override else None), remaining
def multigpu_launch(
name: str,
file: str,
num_gpus: Sequence[int],
env_key: str,
inner: Callable[[], int],
kind: str,
pre_launch_fn: Optional[Callable[[List[int]], None]] = None,
timeout: Optional[int] = None,
) -> NoReturn | None:
"""Shared torchrun-based launcher.
See module docstring. `name` is the caller's `__name__`; `file` is its
`__file__`. `env_key` is a unique string per kind (test/benchmark) used to
detect the inside-torchrun state. `inner` returns an exit code.
`pre_launch_fn`, if given, runs once in the outer process *before* any
torchrun child is spawned. It receives the list of world sizes that will
actually be launched (already filtered against the host's GPU count and
any ``--num-gpu`` override). Use it for parallel JIT precompilation so the
on-disk kernel cache is warm by the time the torchrun children import
their kernels.
`timeout`, if given, bounds each per-world-size torchrun invocation (in
seconds). On expiry the child's whole process group is killed and the
launcher exits non-zero. `None` (the default) waits indefinitely.
"""
pid_key = env_key + "_PID"
if env_key in os.environ:
assert pid_key in os.environ
if name != "__main__":
return
rank = int(os.environ["LOCAL_RANK"])
if rank != 0:
sys.stdout = open(os.devnull, "w")
torch.cuda.set_device(rank)
return sys.exit(inner())
assert pid_key not in os.environ
if name != "__main__":
return logger.warning(
f"{file} can not directly run with `pytest`. "
"Use `python` to invoke it, which will internally relaunch it "
"under torchrun for each requested number of GPUs."
)
num_devices = torch.cuda.device_count()
override, forwarded_args = _extract_num_gpus_override(sys.argv[1:])
if override is not None:
logger.info(f"--num-gpu override: running only with {override}")
num_gpus = override
for N in num_gpus:
if N <= 1 or N > num_devices:
raise ValueError(
f"Invalid number of GPUs requested: {N} "
f"(available: {num_devices})"
)
os.environ[env_key] = "1"
os.environ[pid_key] = str(os.getpid())
os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("GLOO_SOCKET_IFNAME", "lo") # single-machine setup
signal.signal(signal.SIGINT, signal.default_int_handler)
runnable: List[int] = []
for N in sorted(num_gpus):
assert N > 1
if N > num_devices:
logger.warning(f"Skipping {kind} with {N} GPUs ({num_devices} available)")
continue
runnable.append(N)
if pre_launch_fn is not None and runnable:
logger.info(f"Running pre-launch hook for world sizes {runnable}")
pre_launch_fn(runnable)
for N in runnable:
logger.info(f"Running {kind} with {N} GPUs")
cmd = [
"torchrun",
"--nproc_per_node",
str(N),
"--local-addr",
"127.0.0.1",
file,
]
cmd += forwarded_args
proc = subprocess.Popen(cmd, start_new_session=True)
pgid = proc.pid
returncode = -1
timed_out = False
try:
returncode = proc.wait(timeout=timeout)
except subprocess.TimeoutExpired:
timed_out = True
finally:
_kill_descendants(os.getpid())
_kill_pgroup(pgid)
try:
proc.wait(timeout=5)
except subprocess.TimeoutExpired:
pass
if timed_out:
logger.error(f"{kind} (nproc={N}) timed out after {timeout} seconds")
sys.exit(1)
if returncode != 0:
logger.error(f"{kind} failed with {N} GPUs (exit {returncode})")
sys.exit(returncode)
logger.info(f"All {kind}s passed")
+48 -40
View File
@@ -1,49 +1,57 @@
import os
import subprocess
import sys
from typing import Callable
from typing import Callable, List, Optional, Sequence
import pytest
from sglang.jit_kernel.mp import multigpu_launch
def multiprocess_test(file: str, nproc: int, timeout: int = 240) -> None:
"""Launch this script as a torchrun worker and assert success.
The default budget covers the cold-cache first invocation, where the
worker pays the full triton + cutlass JIT compile cost (60-180s observed
on H200). The previous 90s default tripped intermittently on the first
parametrisation of `test_tp_qknorm` (seen on `main` runs too, not only
on fresh-venv PRs); subsequent parametrisations finished in ~60s once
the JIT cache was warm.
def multigpu_pytest_main(
name: str,
file: str,
num_gpus: Sequence[int],
*,
pre_launch_fn: Optional[Callable[[List[int]], None]] = None,
timeout: Optional[int] = 600,
) -> None:
"""cudalib-style multi-GPU pytest entry point.
Drop this at the bottom of a test file::
multigpu_pytest_main(__name__, __file__, num_gpus=range(2, 9))
When the file is run with ``python <file>``, it relaunches itself under
``torchrun --nproc_per_node=N <file>`` for each N in ``num_gpus``. Inside
each worker, ``pytest.main([file, ...forwarded_args])`` runs the collected
tests. Pass ``--num-gpu 2,4`` on the command line to override ``num_gpus``.
``pre_launch_fn`` (kw-only) runs once in the outer process before any
torchrun child starts, receiving the runnable world sizes. Use it for
parallel JIT precompilation so torchrun children hit a warm disk cache
instead of compiling kernels on first call.
``timeout`` (kw-only, seconds) bounds each per-world-size torchrun
invocation. The default budget covers the cold-cache first invocation
(the worker pays the full triton + cutlass JIT compile cost, 60-180s
observed on H200) plus the nightly full sweep, which runs every size x
dtype x algo x graph-mode parametrisation rather than the reduced in-CI
range. A worker that exceeds the budget is killed and the run fails. Pass
``None`` to wait indefinitely.
"""
cmd = [
"torchrun",
f"--nproc_per_node={nproc}",
def inner() -> int:
# CI's run_unittest_files invokes `python3 <file> -f` (legacy
# unittest failfast). Translate to pytest's `-x` so it survives.
pytest_args = ["-x" if a == "-f" else a for a in sys.argv[1:]]
return pytest.main([file] + pytest_args)
return multigpu_launch(
name,
file,
]
try:
result = subprocess.run(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
timeout=timeout,
)
except subprocess.TimeoutExpired as e:
raise RuntimeError(
f"torchrun (nproc={nproc}) timed out after {timeout} seconds\n"
f"{e.stdout}"
) from e
assert result.returncode == 0, (
f"torchrun (nproc={nproc}) failed with rc={result.returncode}\n"
f"{result.stdout}"
num_gpus,
env_key="_IS_TEST_MULTIGPU_SGLANG_JIT_KERNEL",
inner=inner,
kind="test",
pre_launch_fn=pre_launch_fn,
timeout=timeout,
)
def multiprocess_main(file: str, main: Callable[[], None]) -> None:
"""Helper to run a function in a multiprocess torchrun context."""
if "LOCAL_RANK" in os.environ:
main()
else:
sys.exit(pytest.main([file, "-v", "-s"]))
@@ -1,27 +1,34 @@
"""
Benchmark JIT custom all-reduce (v2) vs NCCL vs AOT custom all-reduce (v1).
"""Benchmark JIT custom all-reduce (v2) vs NCCL, AOT custom-AR (v1), and
FlashInfer trtllm allreduce_fusion.
Usage (torchrun required for multi-GPU):
torchrun --nproc_per_node=2 bench_custom_all_reduce.py
torchrun --nproc_per_node=4 bench_custom_all_reduce.py --dtype float16
torchrun --nproc_per_node=8 bench_custom_all_reduce.py --warmup 10 --iters 100
Usage::
The script initializes all three backends, then benchmarks each over a sweep
of message sizes. Results are printed as a comparison table on rank 0.
# Benchmark on every supported world size (2..8 GPUs):
python benchmark/bench_custom_all_reduce.py
# Pick a specific world size (or comma-separated list):
python benchmark/bench_custom_all_reduce.py --num-gpu 4
python benchmark/bench_custom_all_reduce.py --num-gpu 2,4,8
The script self-relaunches under ``torchrun --nproc_per_node=N`` for each N in
``num_gpus``; results are printed on rank 0 of every run.
"""
import argparse
from __future__ import annotations
import atexit
import contextlib
import gc
import logging
import os
from math import isnan
from typing import Dict, List, Optional
from typing import Optional
import torch
import torch.distributed as dist
from sglang.jit_kernel.benchmark.utils import is_in_ci
import sglang.srt.distributed.parallel_state as ps
from sglang.jit_kernel.benchmark import marker
from sglang.jit_kernel.benchmark.utils import get_benchmark_range, multigpu_bench_main
from sglang.jit_kernel.mp import register_comm_cleanup
from sglang.jit_kernel.utils import cache_once, is_arch_support_pdl
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -30,12 +37,14 @@ register_cuda_ci(
disabled="requires multi-GPU, self-skips in CI",
)
DTYPE_MAP = {
"float16": torch.float16,
"bfloat16": torch.bfloat16,
"float32": torch.float32,
}
# ---------------------------------------------------------------------------
# Sweep parameters
# ---------------------------------------------------------------------------
DTYPE = torch.bfloat16
# torch.dtype.itemsize exists only on newer torch; element_size() is portable.
DTYPE_ITEMSIZE = torch.tensor([], dtype=DTYPE).element_size()
MESSAGE_SIZES_BYTES = [
4 * 1024, # 4K
16 * 1024, # 16K
@@ -50,29 +59,65 @@ MESSAGE_SIZES_BYTES = [
7 * 128 * 1024, # 896K
1 * 1024 * 1024, # 1M
2 * 1024 * 1024, # 2M
3 * 1024 * 1024, # 2M
3 * 1024 * 1024, # 3M
4 * 1024 * 1024, # 4M
8 * 1024 * 1024, # 8M
16 * 1024 * 1024, # 16M
32 * 1024 * 1024, # 32M
]
WORLD_SIZES = list(range(2, 9))
MAX_BYTES = max(MESSAGE_SIZES_BYTES)
# trtllm allreduce_fusion only supports these world sizes.
FI_SUPPORTED_WORLD_SIZES = (2, 4, 8)
# AOT custom_all_reduce (v1) only supports these world sizes.
AOT_SUPPORTED_WORLD_SIZES = (2, 4, 6, 8)
PROVIDERS = ["nccl", "aot", "jit", "fi"]
WORLD_SIZES = get_benchmark_range(WORLD_SIZES, [2, 4, 8])
# ---------------------------------------------------------------------------
# Per-rank distributed init (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
# Quieter benchmark output.
logging.disable(logging.INFO)
torch.cuda.set_stream(torch.cuda.Stream())
return coord.cpu_group
@cache_once
def _init_nccl_group() -> dist.ProcessGroup:
_init_cpu_group()
coord = ps._WORLD
assert coord is not None and coord.device_group is not None
return coord.device_group
# ---------------------------------------------------------------------------
# Backend wrappers - each exposes a uniform interface:
# .name - display name
# .capture() - context manager for CUDA-graph recording
# .all_reduce() - perform an all-reduce and return the result tensor
# Backend wrappers - each exposes:
# .all_reduce(tensor) -> Tensor
# .graph_context() -> context manager wrapping cuda-graph capture
# (nullcontext when capture is not required)
# ---------------------------------------------------------------------------
class NCCLAllReduceBackend:
name = "NCCL"
def __init__(self) -> None:
self.group = _init_nccl_group()
def __init__(self, group: dist.ProcessGroup):
self.group = group
def capture(self, register_input: bool):
def graph_context(self):
return contextlib.nullcontext()
def all_reduce(self, tensor: torch.Tensor) -> torch.Tensor:
@@ -80,42 +125,42 @@ class NCCLAllReduceBackend:
return tensor
class AOTAllReduceBackend:
name = "AOT"
def __init__(self, group: dist.ProcessGroup, device: torch.device):
from sglang.srt.distributed.device_communicators.custom_all_reduce import (
CustomAllreduce,
class JITAllReduceBackend:
def __init__(self) -> None:
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
max_size = max(MESSAGE_SIZES_BYTES)
self.comm = CustomAllreduce(group, device, max_size=max_size)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
self.comm = CustomAllReduceV2(
_init_cpu_group(), device, max_pull_size=MAX_BYTES
)
if self.comm.disabled:
raise RuntimeError("AOT CustomAllreduce is disabled on this system")
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
register_comm_cleanup(self.comm)
def capture(self, register_input: bool):
return self.comm.capture() # ignore register_input since v1 always requires it
def graph_context(self):
return self.comm.capture()
def all_reduce(self, tensor: torch.Tensor) -> Optional[torch.Tensor]:
assert self.comm.should_custom_ar(tensor), str(tensor.shape)
return self.comm.custom_all_reduce(tensor)
class JITAllReduceBackend:
name = "JIT"
def __init__(self, group: dist.ProcessGroup, device: torch.device):
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
class AOTAllReduceBackend:
def __init__(self) -> None:
from sglang.srt.distributed.device_communicators.custom_all_reduce import (
CustomAllreduce,
)
max_size = max(MESSAGE_SIZES_BYTES)
self.comm = CustomAllReduceV2(group, device, max_pull_size=max_size)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
self.comm = CustomAllreduce(_init_cpu_group(), device, max_size=MAX_BYTES)
if self.comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
raise RuntimeError("AOT CustomAllreduce is disabled on this system")
register_comm_cleanup(self.comm)
def capture(self, register_input: bool):
return self.comm.capture() if register_input else contextlib.nullcontext()
def graph_context(self):
return self.comm.capture()
def all_reduce(self, tensor: torch.Tensor) -> Optional[torch.Tensor]:
assert self.comm.should_custom_ar(tensor), str(tensor.shape)
@@ -123,262 +168,127 @@ class JITAllReduceBackend:
class FlashInferAllReduceBackend:
name = "FI"
def __init__(self, group: dist.ProcessGroup, dtype: torch.dtype):
def __init__(self) -> None:
import flashinfer.comm as comm
rank = torch.distributed.get_rank(group=group)
world_size = torch.distributed.get_world_size(group=group)
max_size = max(MESSAGE_SIZES_BYTES)
hidden_dim = min(MESSAGE_SIZES_BYTES) // 2
num_tokens = max_size // hidden_dim
self.comm = comm
self.hidden_dim = hidden_dim
self.workspace = comm.create_allreduce_fusion_workspace(
group = _init_cpu_group()
rank = dist.get_rank(group=group)
world_size = dist.get_world_size(group=group)
# Use the smallest message size as the inner hidden dim, so any
# message in the sweep is an integer multiple of it.
hidden_dim = min(MESSAGE_SIZES_BYTES) // DTYPE_ITEMSIZE
num_tokens = MAX_BYTES // (hidden_dim * DTYPE_ITEMSIZE)
self._comm = comm
self._hidden_dim = hidden_dim
self._workspace = comm.create_allreduce_fusion_workspace(
backend="trtllm",
world_size=world_size,
rank=rank,
max_token_num=num_tokens,
hidden_dim=hidden_dim,
dtype=dtype,
dtype=DTYPE,
)
def capture(self, *_):
def graph_context(self):
return contextlib.nullcontext()
def all_reduce(self, tensor: torch.Tensor) -> Optional[torch.Tensor]:
return self.comm.allreduce_fusion(
input=tensor.view(-1, self.hidden_dim),
workspace=self.workspace,
pattern=self.comm.AllReduceFusionPattern.kAllReduce,
launch_with_pdl=True,
def all_reduce(self, tensor: torch.Tensor) -> torch.Tensor:
return self._comm.allreduce_fusion(
input=tensor.view(-1, self._hidden_dim),
workspace=self._workspace,
pattern=self._comm.AllReduceFusionPattern.kAllReduce,
launch_with_pdl=is_arch_support_pdl(),
fp32_acc=True,
)
# ---------------------------------------------------------------------------
# Benchmarking helpers
# ---------------------------------------------------------------------------
@cache_once
def _init_nccl_backend() -> NCCLAllReduceBackend:
return NCCLAllReduceBackend()
def parse_args():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--dtype", choices=DTYPE_MAP.keys(), default="bfloat16")
p.add_argument("--warmup", type=int, default=5)
p.add_argument("--iters", type=int, default=50)
p.add_argument("--no-inplace", dest="register_input", action="store_false")
return p.parse_args()
@cache_once
def _init_jit_backend() -> JITAllReduceBackend:
return JITAllReduceBackend()
@torch.inference_mode()
def bench_one(
backend,
inp: torch.Tensor,
warmup: int,
iters: int,
group: dist.ProcessGroup,
register_input: bool,
) -> float:
@cache_once
def _init_aot_backend() -> AOTAllReduceBackend:
return AOTAllReduceBackend()
@cache_once
def _init_fi_backend() -> FlashInferAllReduceBackend:
return FlashInferAllReduceBackend()
BACKEND_FACTORY = {
"nccl": _init_nccl_backend,
"jit": _init_jit_backend,
"aot": _init_aot_backend,
"fi": _init_fi_backend,
}
@cache_once
def _init_all_backends() -> None:
"""Pre-build every supported backend before any timed iteration so JIT
compilation / IPC setup don't bleed into the first measured size.
"""
Run *warmup* iterations of all-reduce first.
Return the average time for *iters* iterations of all-reduce.
"""
dist.barrier(group=group)
for _ in range(warmup):
backend.all_reduce(inp)
torch.cuda.synchronize()
# Capture a CUDA graph with *iters* all-reduce calls.
inp_batch = torch.stack([inp] * 4)
graph = torch.cuda.CUDAGraph()
with backend.capture(register_input):
with torch.cuda.graph(graph):
for i in range(iters):
backend.all_reduce(inp_batch[i % 4])
torch.cuda.synchronize()
# Warm up the graph once.
graph.replay()
# Timed replay.
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
torch.cuda.synchronize()
dist.barrier(group=group)
graph.replay() # make the stream busy
start.record()
graph.replay()
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) / iters
def bench_sweep(
backend,
sizes_bytes: List[int],
dtype: torch.dtype,
device: torch.device,
warmup: int,
iters: int,
group: dist.ProcessGroup,
register_input: bool,
) -> Dict[int, float]:
"""Benchmark one backend over all message sizes."""
elem_size = torch.tensor([], dtype=dtype).element_size()
results: Dict[int, float] = {}
for sz in sizes_bytes:
numel = sz // elem_size
inp = torch.zeros(numel, dtype=dtype, device=device)
try:
elapsed_ms = bench_one(backend, inp, warmup, iters, group, register_input)
results[sz] = elapsed_ms * 1000 # convert to us per iter
except AssertionError:
results[sz] = float("nan")
return results
world_size = dist.get_world_size(_init_cpu_group())
factories = dict(BACKEND_FACTORY)
if world_size not in AOT_SUPPORTED_WORLD_SIZES:
factories.pop("aot")
if world_size not in FI_SUPPORTED_WORLD_SIZES:
factories.pop("fi")
for fn in factories.values():
fn()
# ---------------------------------------------------------------------------
# Result printing
# Benchmark
# ---------------------------------------------------------------------------
def print_results(
backends: list,
all_results: Dict[str, Dict[int, float]],
sizes_bytes: List[int],
) -> None:
"""Print a comparison table on rank 0."""
def human_bytes(n: int) -> str:
for suffix, unit in [("M", 1 << 20), ("K", 1 << 10)]:
if n >= unit and n % unit == 0:
return f"{n // unit}{suffix}"
return f"{n}B"
def fmt_us(v: float) -> str:
return f"{v:13.1f}" if not isnan(v) else " n/a"
names = [b.name for b in backends]
nccl_name = "NCCL"
# Header
header_cols = [f"{n:>13}" for n in names]
speedup_cols = [f"{n:>13}/NCCL" for n in names if n != nccl_name]
header = f"{'Size':>8} " + " ".join(header_cols)
for sc in speedup_cols:
header += f" {sc}"
header += " "
print()
print(header)
print("-" * len(header))
# Rows
for sz in sizes_bytes:
row = f"{human_bytes(sz):>8}"
nccl_lat = all_results[nccl_name][sz]
for n in names:
row += f" {fmt_us(all_results[n][sz])}"
for n in names:
if n == nccl_name:
continue
lat = all_results[n][sz]
if not isnan(lat):
row += f" {nccl_lat / lat:17.2f}x"
else:
row += f" {'n/a':>17}"
print(row)
# ---------------------------------------------------------------------------
# Distributed setup
# ---------------------------------------------------------------------------
def init_distributed():
"""Initialize distributed groups using torchrun env vars.
Returns (rank, world_size, device, cpu_group, nccl_group).
"""
import sglang.srt.distributed.parallel_state as ps
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
world_size = int(os.environ.get("WORLD_SIZE", "1"))
rank = local_rank
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)
torch.cuda.set_stream(torch.cuda.Stream()) # use a non-default stream
torch.distributed.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
@marker.parametrize("message_bytes", MESSAGE_SIZES_BYTES)
@marker.benchmark("provider", PROVIDERS)
def benchmark(message_bytes: int, provider: str):
cpu_group = _init_cpu_group()
gpu_group = _init_nccl_group()
world_size = dist.get_world_size(cpu_group)
if provider == "fi" and world_size not in FI_SUPPORTED_WORLD_SIZES:
marker.skip(
f"flashinfer trtllm allreduce_fusion needs world_size in "
f"{FI_SUPPORTED_WORLD_SIZES}"
)
if provider == "aot" and world_size not in AOT_SUPPORTED_WORLD_SIZES:
marker.skip(
f"AOT custom_all_reduce needs world_size in " f"{AOT_SUPPORTED_WORLD_SIZES}"
)
_init_all_backends()
backend = BACKEND_FACTORY[provider]()
numel = message_bytes // DTYPE_ITEMSIZE
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
x = torch.randn(numel, dtype=DTYPE, device=device)
# Bandwidth-equivalent bytes moved by a ring all-reduce per rank.
effective_bytes = int(x.nbytes * 2 * (world_size - 1) / world_size)
return marker.do_bench(
backend.all_reduce,
input_args=(x,),
graph_context_fn=backend.graph_context,
sync_multigpu_fn=lambda: dist.barrier(gpu_group),
# all-reduce is in-place w.r.t. its argument; explicit footprint
# captures the cross-GPU traffic instead.
memory_args=None,
memory_output=None,
extra_memory_footprint=effective_bytes,
)
cpu_group = coord.cpu_group
nccl_group = coord.device_group
assert nccl_group is not None
return rank, world_size, device, cpu_group, nccl_group
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
logging.basicConfig(level=logging.WARNING)
args = parse_args()
dtype = DTYPE_MAP[args.dtype]
rank, world_size, device, cpu_group, nccl_group = init_distributed()
# Instantiate backends.
backends = [
NCCLAllReduceBackend(nccl_group),
JITAllReduceBackend(cpu_group, device),
]
if world_size in [2, 4, 6, 8]:
backends.insert(1, AOTAllReduceBackend(cpu_group, device))
if world_size in [2, 4, 8]:
backends.append(FlashInferAllReduceBackend(cpu_group, dtype))
# Run benchmarks.
all_results: Dict[str, Dict[int, float]] = {}
torch.cuda.synchronize()
for backend in backends:
if rank == 0:
print(f"Benchmarking {backend.name} ...")
all_results[backend.name] = bench_sweep(
backend,
MESSAGE_SIZES_BYTES,
dtype,
device,
args.warmup,
args.iters,
cpu_group,
args.register_input,
)
# Aggregate across ranks (use max to reflect the slowest rank).
for name in list(all_results):
for sz in MESSAGE_SIZES_BYTES:
val = all_results[name].get(sz)
if val is None:
continue
t = torch.tensor([val], dtype=torch.float64, device=device)
dist.all_reduce(t, op=dist.ReduceOp.MAX, group=nccl_group)
all_results[name][sz] = t.item()
# Print results on rank 0.
if rank == 0:
print_results(backends, all_results, MESSAGE_SIZES_BYTES)
del backends, all_results
gc.collect()
dist.destroy_process_group()
if __name__ == "__main__" and not is_in_ci():
main()
if __name__ == "__main__":
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=WORLD_SIZES,
main_fn=benchmark.run,
)
+158 -94
View File
@@ -1,17 +1,39 @@
"""Benchmark fused TP QKNorm (push-mode custom-AR + RMSNorm) vs the serial
baseline (RMS sum-sq -> pull-mode all-reduce -> RMS apply).
Usage::
# Benchmark on every supported world size (2..8 GPUs):
python benchmark/bench_tp_qknorm.py
# Specific world sizes:
python benchmark/bench_tp_qknorm.py --num-gpu 4
python benchmark/bench_tp_qknorm.py --num-gpu 2,4,8
"""
from __future__ import annotations
import argparse
import atexit
import logging
import multiprocessing
import os
from multiprocessing.context import SpawnProcess
from typing import List
import torch
import torch.distributed as dist
import sglang.srt.distributed.parallel_state as ps
from sglang.jit_kernel.all_reduce import (
_jit_custom_all_reduce_pull_module,
_jit_custom_all_reduce_push_module,
_jit_fused_parallel_qknorm_module,
fused_parallel_qknorm,
get_fused_parallel_qknorm_max_occupancy,
)
from sglang.jit_kernel.utils import get_ci_test_range
from sglang.jit_kernel.benchmark import marker
from sglang.jit_kernel.benchmark.utils import multigpu_bench_main
from sglang.jit_kernel.mp import register_comm_cleanup
from sglang.jit_kernel.utils import cache_once, get_ci_test_range
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
@@ -23,80 +45,131 @@ register_cuda_ci(
disabled="requires multi-GPU, self-skips in CI",
)
Q_K_DIMS = [(6144, 1024)]
# ---------------------------------------------------------------------------
# Sweep parameters
# ---------------------------------------------------------------------------
DTYPE = torch.bfloat16
EPS = 1e-6
Q_K_DIMS = [(6144, 1024)]
BATCH_SIZES = get_ci_test_range([2**i for i in range(15)], [1, 64, 1024])
NUM_LAYERS = 8
MAX_PUSH_SIZE = 8 * max(BATCH_SIZES)
PROVIDERS = ["fused", "baseline"]
def parse_args():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--iters", type=int, default=100)
return parser.parse_args()
# ---------------------------------------------------------------------------
# Parallel JIT precompile (outer process, before any torchrun child starts)
# ---------------------------------------------------------------------------
def init_distributed():
def _compile_one(world_size: int) -> None:
"""Compile every kernel this bench touches for a single world_size.
Top-level so it survives ``spawn`` pickling. Compiled artifacts are
cached on disk by ``tvm_ffi``; torchrun children will reuse them.
"""
# baseline path: sum-sq -> pull-mode all-reduce -> apply
_jit_custom_all_reduce_pull_module(DTYPE, world_size)
# fused path: push-mode all-reduce
_jit_custom_all_reduce_push_module(DTYPE, world_size)
# fused path: fused QKNorm kernel (one per (dtype, world_size, q_dim, k_dim))
for q_dim, k_dim in Q_K_DIMS:
_jit_fused_parallel_qknorm_module(DTYPE, world_size, q_dim, k_dim)
def _precompile_kernels(num_gpus: List[int]) -> None:
ctx = multiprocessing.get_context("spawn")
procs: list[tuple[int, SpawnProcess]] = []
for world_size in num_gpus:
p = ctx.Process(target=_compile_one, args=(world_size,))
p.start()
procs.append((world_size, p))
for world_size, p in procs:
p.join()
if p.exitcode != 0:
raise RuntimeError(
f"TP QKNorm precompile failed for {world_size=} " f"(exit {p.exitcode})"
)
# ---------------------------------------------------------------------------
# Per-rank distributed init (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
rank = local_rank
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
logging.disable(logging.INFO)
torch.cuda.set_stream(torch.cuda.Stream())
return coord.cpu_group
cpu_group = coord.cpu_group
@cache_once
def _init_gpu_group() -> dist.ProcessGroup:
_init_cpu_group()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
gpu_group = dist.new_group(backend="nccl", device_id=device)
assert isinstance(gpu_group, dist.ProcessGroup)
atexit.register(lambda: dist.destroy_process_group(gpu_group))
return gpu_group
@cache_once
def _init_fused_comm() -> CustomAllReduceV2:
"""Push-mode workspace sized for the fused-QKNorm bench."""
cpu_group = _init_cpu_group()
world_size = dist.get_world_size(cpu_group)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
q_dim, k_dim = Q_K_DIMS[0]
max_occupancy = get_fused_parallel_qknorm_max_occupancy(
DTYPE, world_size, Q_K_DIMS[0][0], Q_K_DIMS[0][1]
DTYPE, world_size, q_dim, k_dim
)
if rank == 0:
if dist.get_rank(cpu_group) == 0:
print(f"Max occupancy for fused_parallel_qknorm: {max_occupancy} blocks/SM")
props = torch.cuda.get_device_properties(device)
comm = CustomAllReduceV2(
cpu_group,
device,
max_pull_size=0,
max_push_size=8 * max(BATCH_SIZES),
max_push_size=MAX_PUSH_SIZE,
max_push_blocks=props.multi_processor_count * max_occupancy,
)
comm_ = CustomAllReduceV2(cpu_group, device)
if comm.disabled or comm_.disabled:
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
return rank, world_size, device, cpu_group, comm, comm_
register_comm_cleanup(comm)
return comm
@torch.inference_mode()
def bench_one(fn, warmup: int, iters: int) -> float:
for _ in range(warmup):
fn(0)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
for i in range(NUM_LAYERS):
fn(i)
graph.replay()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
graph.replay()
start.record()
for i in range(iters):
graph.replay()
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) * 1000.0 / (iters * NUM_LAYERS)
@cache_once
def _init_baseline_comm() -> CustomAllReduceV2:
"""Default (pull-mode) workspace for the serial baseline."""
cpu_group = _init_cpu_group()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
comm = CustomAllReduceV2(cpu_group, device)
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
register_comm_cleanup(comm)
return comm
def rmsnorm_baseline(
comm_,
# ---------------------------------------------------------------------------
# Implementations
# ---------------------------------------------------------------------------
def _rmsnorm_baseline(
comm: CustomAllReduceV2,
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
@@ -106,65 +179,56 @@ def rmsnorm_baseline(
from sglang.srt.models.minimax_m2 import rms_apply_serial, rms_sumsq_serial
sum_sq = rms_sumsq_serial(q, k)
sum_sq = comm_.custom_all_reduce(sum_sq)
sum_sq = comm.custom_all_reduce(sum_sq)
rms_apply_serial(q, k, q_weight, k_weight, sum_sq, world_size, EPS)
def main():
args = parse_args()
rank, world_size, device, _, comm, comm_ = init_distributed()
torch.cuda.set_stream(torch.cuda.Stream())
# ---------------------------------------------------------------------------
# Benchmark
# ---------------------------------------------------------------------------
if rank == 0:
print(
f"{'q_dim':>8} {'k_dim':>8} {'batch':>8} {'fused_us':>12} {'baseline_us':>12}"
)
for q_dim, k_dim in Q_K_DIMS:
local_q_dim = q_dim // world_size
local_k_dim = k_dim // world_size
for batch_size in BATCH_SIZES:
q = torch.randn(
NUM_LAYERS, batch_size, local_q_dim, device=device, dtype=DTYPE
)
k = torch.randn(
NUM_LAYERS, batch_size, local_k_dim, device=device, dtype=DTYPE
)
q_weight = torch.randn(NUM_LAYERS, local_q_dim, device=device, dtype=DTYPE)
k_weight = torch.randn(NUM_LAYERS, local_k_dim, device=device, dtype=DTYPE)
@marker.parametrize("q_dim,k_dim", Q_K_DIMS)
@marker.parametrize("batch_size", BATCH_SIZES)
@marker.benchmark("provider", PROVIDERS)
def benchmark(q_dim: int, k_dim: int, batch_size: int, provider: str):
cpu_group = _init_cpu_group()
gpu_group = _init_gpu_group()
world_size = dist.get_world_size(cpu_group)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
local_q_dim = q_dim // world_size
local_k_dim = k_dim // world_size
def run_fused(i: int):
fused_parallel_qknorm(
comm.obj,
q[i],
k[i],
q_weight[i],
k_weight[i],
EPS,
)
q = torch.randn(batch_size, local_q_dim, device=device, dtype=DTYPE)
k = torch.randn(batch_size, local_k_dim, device=device, dtype=DTYPE)
q_weight = torch.randn(local_q_dim, device=device, dtype=DTYPE)
k_weight = torch.randn(local_k_dim, device=device, dtype=DTYPE)
def run_baseline(i: int):
rmsnorm_baseline(
comm_,
q[i],
k[i],
q_weight[i],
k_weight[i],
world_size,
)
if provider == "fused":
comm = _init_fused_comm()
fused_us = bench_one(run_fused, args.warmup, args.iters)
baseline_us = bench_one(run_baseline, args.warmup, args.iters)
def fn(q, k, q_weight, k_weight):
fused_parallel_qknorm(comm.obj, q, k, q_weight, k_weight, EPS)
if rank == 0:
print(
f"{q_dim:8d} {k_dim:8d} {batch_size:8d} "
f"{fused_us:12.1f} {baseline_us:12.1f}"
)
else:
comm = _init_baseline_comm()
comm.close()
dist.destroy_process_group()
def fn(q, k, q_weight, k_weight):
_rmsnorm_baseline(comm, q, k, q_weight, k_weight, world_size)
return marker.do_bench(
fn,
input_args=(q, k, q_weight, k_weight),
sync_multigpu_fn=lambda: dist.barrier(gpu_group),
memory_output=(q, k), # NOTE: In-place updates on q, k;
)
if __name__ == "__main__":
main()
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=[2, 4, 8], # NOTE: don't support other world size now
main_fn=benchmark.run,
pre_launch_fn=_precompile_kernels,
)
+130 -131
View File
@@ -1,28 +1,34 @@
"""
Correctness test for the JIT custom all-reduce (v2) kernel.
"""Correctness test for the JIT custom all-reduce (v2) kernel.
The test compares the JIT custom all-reduce output against NCCL all-reduce
for various tensor sizes and dtypes, in both eager and CUDA-graph modes.
Compares the JIT custom all-reduce output against NCCL all-reduce for a sweep
of tensor sizes, dtypes, and algorithms, in both eager and CUDA-graph modes.
Usage:
python -m pytest test_jit_custom_all_reduce.py -v
Usage::
This file doubles as the torchrun worker script. The test class launches
torchrun --nproc_per_node=N <this_file>
and asserts that all worker processes exit successfully.
# Run the test on the default world sizes (2, 4, 8 GPUs):
python tests/test_custom_all_reduce.py
# Pick a specific world size (or comma-separated list), e.g. the rarer
# odd / non-power-of-two counts that the default sweep skips:
python tests/test_custom_all_reduce.py --num-gpu 3
python tests/test_custom_all_reduce.py --num-gpu 2,4,6,8
# Extra pytest args (forwarded to each torchrun worker):
python tests/test_custom_all_reduce.py -k bfloat16
"""
from __future__ import annotations
import atexit
import itertools
import logging
import multiprocessing as mp
import multiprocessing
import os
from typing import Dict, Optional, Tuple
from multiprocessing.context import SpawnProcess
from typing import List
import pytest
import torch
import torch.distributed as dist
import triton
import sglang.srt.distributed.parallel_state as ps
from sglang.jit_kernel.all_reduce import (
@@ -30,7 +36,9 @@ from sglang.jit_kernel.all_reduce import (
_jit_custom_all_reduce_pull_module,
_jit_custom_all_reduce_push_module,
)
from sglang.jit_kernel.tests.utils import multiprocess_main, multiprocess_test
from sglang.jit_kernel.mp import register_comm_cleanup
from sglang.jit_kernel.tests.utils import multigpu_pytest_main
from sglang.jit_kernel.utils import cache_once, get_ci_test_range
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
@@ -47,7 +55,7 @@ register_cuda_ci(
)
# ---------------------------------------------------------------------------
# Test parameters (shared between test class and worker)
# Test parameters
# ---------------------------------------------------------------------------
TEST_SIZES = [
@@ -59,181 +67,172 @@ TEST_SIZES = [
4 * 1024,
32 * 1024,
256 * 1024,
2 * 1024 * 1024, # 2M elements
4 * 1024 * 1024, # 4M elements
2 * 1024 * 1024,
4 * 1024 * 1024,
]
TEST_DTYPES = [torch.float16, torch.bfloat16, torch.float32]
SHOTS = [
TEST_ALGOS = [
AllReduceAlgo.ONE_SHOT_PULL,
AllReduceAlgo.ONE_SHOT_PUSH,
AllReduceAlgo.TWO_SHOT_PULL,
]
USE_GRAPH_OPTIONS = [True, False]
TEST_CONFIG = itertools.product(TEST_SIZES, TEST_DTYPES, SHOTS, USE_GRAPH_OPTIONS)
USE_GRAPH_OPTIONS = [False, True]
TEST_LAYERS = 4
TEST_LOOP = 16
TEST_SIZES = get_ci_test_range(TEST_SIZES, [16, 1024, 32 * 1024, 2 * 1024 * 1024])
TEST_DTYPES = get_ci_test_range(TEST_DTYPES, [torch.bfloat16])
# ---------------------------------------------------------------------------
# Test class (runs via pytest, launches torchrun subprocesses)
# Parallel JIT precompile (outer process, before any torchrun child starts)
# ---------------------------------------------------------------------------
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 _compile_one(dtype: torch.dtype, world_size: int) -> None:
"""Compile both (push, pull) variants for a single (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:
if nproc == 1: # NOTE: special case to speed up tests
return _precompile_kernels()
device_count = torch.cuda.device_count()
if device_count < nproc:
pytest.skip(
f"Requires at least {nproc} GPUs, but only {device_count} available"
)
multiprocess_test(__file__, nproc)
# ---------------------------------------------------------------------------
# Worker logic (executed by each torchrun process)
# ---------------------------------------------------------------------------
def init_distributed():
"""Initialize distributed groups via torchrun env vars.
Returns (rank, device, cpu_group, nccl_group, comm).
Top-level so it survives ``spawn`` pickling. Compiled artifacts are
cached on disk by ``tvm_ffi``; torchrun children will reuse them.
"""
_jit_custom_all_reduce_pull_module(dtype, world_size)
_jit_custom_all_reduce_push_module(dtype, world_size)
def _precompile_kernels(num_gpus: List[int]) -> None:
"""Fan out one process per (dtype, world_size) to warm the JIT cache.
Without this, every torchrun child serial-compiles its kernels on first
use, multiplying the wall-clock cost of the run by ~(#dtypes * #ranks).
"""
ctx = multiprocessing.get_context("spawn")
procs: list[tuple[torch.dtype, int, SpawnProcess]] = []
for dtype, world_size in itertools.product(TEST_DTYPES, num_gpus):
p = ctx.Process(target=_compile_one, args=(dtype, world_size))
p.start()
procs.append((dtype, world_size, p))
for dtype, world_size, p in procs:
p.join()
if p.exitcode != 0:
raise RuntimeError(
f"Custom-all-reduce precompile failed for "
f"{dtype=} {world_size=} (exit {p.exitcode})"
)
# ---------------------------------------------------------------------------
# Per-rank distributed setup (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group_once() -> dist.ProcessGroup:
"""Initialize gloo world group + cuda device for this rank."""
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
rank = local_rank
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
cpu_group = coord.cpu_group
nccl_group = coord.device_group
assert nccl_group is not None
assert isinstance(cpu_group, dist.ProcessGroup)
# Suppress chatty internal logging for cleaner test output.
logging.disable(logging.INFO)
# Use a non-default stream (mirrors prior behavior).
torch.cuda.set_stream(torch.cuda.Stream())
return cpu_group
max_size = max(TEST_SIZES) * 4
@cache_once
def _init_nccl_group_once() -> dist.ProcessGroup:
_init_cpu_group_once()
coord = ps._WORLD
assert coord is not None and coord.device_group is not None
return coord.device_group
@cache_once
def _init_comm_once() -> CustomAllReduceV2:
cpu_group = _init_cpu_group_once()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
max_size = max(TEST_SIZES) * max(
torch.tensor([], dtype=d).element_size() for d in TEST_DTYPES
)
comm = CustomAllReduceV2(cpu_group, device, max_size, max_size)
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
return rank, device, cpu_group, nccl_group, comm
register_comm_cleanup(comm)
return comm
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("use_graph", USE_GRAPH_OPTIONS)
@pytest.mark.parametrize("algo", TEST_ALGOS)
@pytest.mark.parametrize("dtype", TEST_DTYPES)
@pytest.mark.parametrize("size", TEST_SIZES)
@torch.inference_mode()
def worker_test(
device: torch.device,
nccl_group: dist.ProcessGroup,
comm: CustomAllReduceV2,
def test_custom_all_reduce(
size: int,
dtype: torch.dtype,
use_graph: bool,
algo: AllReduceAlgo,
) -> Optional[RuntimeError]:
use_graph: bool,
) -> None:
nccl_group = _init_nccl_group_once()
comm = _init_comm_once()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
comm.override_algo = algo
def get_run_graph_fn():
if use_graph:
graph = torch.cuda.CUDAGraph()
graph_inp = torch.zeros((TEST_LAYERS, size), dtype=dtype, device=device)
out_jits = []
outs: list[torch.Tensor] = []
with comm.capture():
with torch.cuda.graph(graph):
for i in range(TEST_LAYERS):
out_jits.append(comm.custom_all_reduce(graph_inp[i]))
out_jit = torch.stack(out_jits)
outs.append(comm.custom_all_reduce(graph_inp[i]))
out_jit_stack = torch.stack(outs)
torch.cuda.synchronize()
def run_graph(x: torch.Tensor) -> torch.Tensor:
def run(x: torch.Tensor) -> torch.Tensor:
graph_inp.copy_(x)
graph.replay()
return out_jit.clone()
return out_jit_stack.clone()
return run_graph
else:
def get_run_eager_fn():
def run_eager(x: torch.Tensor) -> torch.Tensor:
def run(x: torch.Tensor) -> torch.Tensor:
eager_inp = x.clone()
out_eagers = []
outs = []
for i in range(TEST_LAYERS):
out_eagers.append(comm.custom_all_reduce(eager_inp[i]))
outs.append(comm.custom_all_reduce(eager_inp[i]))
torch.cuda.synchronize()
return torch.stack(out_eagers)
return torch.stack(outs)
return run_eager
run_fn = get_run_graph_fn() if use_graph else get_run_eager_fn()
num_errors = 0
for _ in range(TEST_LOOP):
# NOTE: 15 * 8 < 128, which is the precision limit for bf16
inp = torch.randint(0, 16, (TEST_LAYERS, size), dtype=dtype, device=device)
assert comm.should_custom_ar(inp[0])
out_ref = inp.clone()
dist.all_reduce(out_ref, group=nccl_group)
out_jit = run_fn(inp)
num_errors += not torch.all(out_jit == out_ref)
if num_errors > 0:
return RuntimeError(
f"Test failed for {size=}, {dtype=}, {algo=}, "
f"{use_graph=} with {num_errors} errors. "
)
return None
def worker_main() -> None:
"""Entry point for each torchrun worker process."""
rank, device, cpu_group, nccl_group, comm = init_distributed()
torch.cuda.set_stream(torch.cuda.Stream())
logging.disable(logging.INFO) # Suppress internal logging for cleaner test output
items = list(enumerate(TEST_CONFIG))
for i, (size, dtype, algo, use_graph) in items:
error = worker_test(device, nccl_group, comm, size, dtype, use_graph, algo)
if error is not None:
print(
f"Worker {rank} failed for {size=}, {dtype=}, "
f"{algo=}, {use_graph=}, iteration={i}\n"
f"Error: {error}"
)
# communicate the result to rank 0 for logging
result = torch.tensor([int(error is not None)])
dist.all_reduce(result, group=cpu_group)
failed = bool(result.item())
if failed:
raise RuntimeError(
f"Test failed on rank {rank} for config: "
f"{size=}, {dtype=}, {algo=}, {use_graph=}"
)
comm.close()
dist.destroy_process_group()
out_jit = run(inp)
# Exact equality, since values are small integers within bf16 precision.
triton.testing.assert_close(out_ref, out_jit, atol=0, rtol=0)
if __name__ == "__main__":
multiprocess_main(__file__, worker_main)
# Only sweep the common world sizes (2, 4, 8) by default: testing every
# count in 2..8 serially overruns the per-file CI time budget, and 3/5/6/7
# are rare in practice. Use --num-gpu to exercise them explicitly.
multigpu_pytest_main(
__name__,
__file__,
num_gpus=(2, 4, 8),
pre_launch_fn=_precompile_kernels,
)
+107 -78
View File
@@ -1,16 +1,30 @@
from __future__ import annotations
import atexit
import itertools
import logging
import multiprocessing
import os
from typing import Optional
from multiprocessing.context import SpawnProcess
from typing import List
import pytest
import torch
import torch.distributed as dist
import triton
from sglang.jit_kernel.all_reduce import fused_parallel_qknorm
from sglang.jit_kernel.tests.utils import multiprocess_main, multiprocess_test
import sglang.srt.distributed.parallel_state as ps
from sglang.jit_kernel.all_reduce import (
_jit_custom_all_reduce_push_module,
_jit_fused_parallel_qknorm_module,
fused_parallel_qknorm,
)
from sglang.jit_kernel.mp import register_comm_cleanup
from sglang.jit_kernel.tests.utils import multigpu_pytest_main
from sglang.jit_kernel.utils import cache_once
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -24,53 +38,96 @@ register_cuda_ci(
)
# ---------------------------------------------------------------------------
# Test parameters
# ---------------------------------------------------------------------------
Q_K_DIMS = [(6144, 1024)]
EPS = 1e-6
BATCH_SIZES = [2**n for n in range(0, 14)]
DTYPES = [torch.float16, torch.bfloat16, torch.float32]
TEST_CONFIG = list(itertools.product(Q_K_DIMS, BATCH_SIZES, DTYPES))
@pytest.mark.parametrize("nproc", [2, 4, 8])
def test_tp_qknorm(nproc: int) -> None:
device_count = torch.cuda.device_count()
if device_count < nproc:
pytest.skip(
f"Requires at least {nproc} GPUs, but only {device_count} available"
)
multiprocess_test(__file__, nproc)
# ---------------------------------------------------------------------------
# Parallel JIT precompile (outer process, before any torchrun child starts)
# ---------------------------------------------------------------------------
def init_distributed():
import sglang.srt.distributed.parallel_state as ps
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
def _compile_one(dtype: torch.dtype, world_size: int) -> None:
"""Compile every kernel this test touches for one (dtype, world_size).
Top-level so it survives ``spawn`` pickling. Compiled artifacts are
cached on disk by ``tvm_ffi``; torchrun children will reuse them.
"""
_jit_custom_all_reduce_push_module(dtype, world_size)
for q_dim, k_dim in Q_K_DIMS:
_jit_fused_parallel_qknorm_module(dtype, world_size, q_dim, k_dim)
def _precompile_kernels(num_gpus: List[int]) -> None:
ctx = multiprocessing.get_context("spawn")
procs: list[tuple[torch.dtype, int, SpawnProcess]] = []
for dtype, world_size in itertools.product(DTYPES, num_gpus):
p = ctx.Process(target=_compile_one, args=(dtype, world_size))
p.start()
procs.append((dtype, world_size, p))
for dtype, world_size, p in procs:
p.join()
if p.exitcode != 0:
raise RuntimeError(
f"TP QKNorm precompile failed for {dtype=} {world_size=} "
f"(exit {p.exitcode})"
)
# ---------------------------------------------------------------------------
# Per-rank distributed setup (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group_once() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
rank = local_rank
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
cpu_group = coord.cpu_group
nccl_group = coord.device_group
assert nccl_group is not None
assert isinstance(cpu_group, dist.ProcessGroup)
logging.disable(logging.INFO)
torch.cuda.set_stream(torch.cuda.Stream())
return cpu_group
@cache_once
def _init_nccl_group_once() -> dist.ProcessGroup:
_init_cpu_group_once()
coord = ps._WORLD
assert coord is not None and coord.device_group is not None
return coord.device_group
@cache_once
def _init_comm_once() -> CustomAllReduceV2:
cpu_group = _init_cpu_group_once()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
max_pull_size = 0
max_push_size = 8 * max(BATCH_SIZES)
comm = CustomAllReduceV2(cpu_group, device, max_pull_size, max_push_size)
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
register_comm_cleanup(comm)
return comm
return rank, world_size, device, cpu_group, nccl_group, comm
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _all_gather_cat(x: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor:
@@ -85,17 +142,26 @@ def _rmsnorm_ref(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Ten
return (x_fp32 * scale * weight.float()).to(x.dtype)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("batch_size", BATCH_SIZES)
@pytest.mark.parametrize("q_k_dim", Q_K_DIMS)
@torch.inference_mode()
def worker_test(
rank: int,
world_size: int,
device: torch.device,
nccl_group: dist.ProcessGroup,
comm,
def test_tp_qknorm(
q_k_dim: tuple[int, int],
batch_size: int,
dtype: torch.dtype,
) -> Optional[RuntimeError]:
) -> None:
nccl_group = _init_nccl_group_once()
comm = _init_comm_once()
rank = dist.get_rank(group=nccl_group)
world_size = dist.get_world_size(group=nccl_group)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
q_dim, k_dim = q_k_dim
local_q_dim = q_dim // world_size
local_k_dim = k_dim // world_size
@@ -115,53 +181,16 @@ def worker_test(
q_expected = q_expected[:, rank * local_q_dim : (rank + 1) * local_q_dim]
k_expected = k_expected[:, rank * local_k_dim : (rank + 1) * local_k_dim]
fused_parallel_qknorm(
comm.obj,
q,
k,
q_weight,
k_weight,
EPS,
)
fused_parallel_qknorm(comm.obj, q, k, q_weight, k_weight, EPS)
try:
triton.testing.assert_close(q, q_expected, atol=1e-2, rtol=1e-2)
triton.testing.assert_close(k, k_expected, atol=1e-2, rtol=1e-2)
except AssertionError as err:
return RuntimeError(
f"TP QKNorm mismatch for {batch_size=}, {dtype=}, {world_size=}, {rank=}: {err}"
)
return None
def worker_main() -> None:
rank, world_size, device, cpu_group, nccl_group, comm = init_distributed()
torch.cuda.set_stream(torch.cuda.Stream())
for q_k_dim, batch_size, dtype in TEST_CONFIG:
error = worker_test(
rank,
world_size,
device,
nccl_group,
comm,
q_k_dim,
batch_size,
dtype,
)
result = torch.tensor([int(error is not None)])
dist.all_reduce(result, group=cpu_group)
if error is not None:
print(str(error))
if bool(result.item()):
raise RuntimeError(
f"TP QKNorm test failed for {q_k_dim=}, {batch_size=}, {dtype=}, {world_size=}"
)
print(f"Rank {rank} passed all tests.")
comm.close()
dist.destroy_process_group()
triton.testing.assert_close(q, q_expected, atol=1e-2, rtol=1e-2)
triton.testing.assert_close(k, k_expected, atol=1e-2, rtol=1e-2)
if __name__ == "__main__":
multiprocess_main(__file__, worker_main)
multigpu_pytest_main(
__name__,
__file__,
num_gpus=(2, 4, 8),
pre_launch_fn=_precompile_kernels,
)