benchmarks/multi_gpu_gemm_rs.py
benchmarks/multi_gpu_gemm_rs.pyBrowse 4 files
2,404 tokens
10,109 bytes
Token encoding: o200k_base
Snapshot da3c07b
← Back to SKILL.md
1# SPDX-License-Identifier: Apache-2.02# SPDX-FileCopyrightText: Copyright contributors to the vLLM project3 4"""Minimal multi-GPU GEMM + reduce-scatter microbenchmark.5 6Run on one node with, for example:7 8 torchrun --standalone --nproc-per-node=8 \9 .agents/skills/kernel-microbenchmark/benchmarks/multi_gpu_gemm_rs.py10"""11 12import argparse13import os14import statistics15from collections.abc import Callable16 17import pandas as pd18import torch19import torch.distributed as dist20import torch.distributed._symmetric_memory as symm_mem21 22from vllm.config import VllmConfig, set_current_vllm_config23from vllm.distributed import cleanup_dist_env_and_memory24from vllm.distributed.parallel_state import (25 get_tp_group,26 init_distributed_environment,27 initialize_model_parallel,28)29 30 31def parse_args() -> argparse.Namespace:32 parser = argparse.ArgumentParser()33 parser.add_argument("--m", type=int, nargs="+", default=[128, 512, 2048])34 parser.add_argument("--n", type=int, default=4096)35 parser.add_argument(36 "--k",37 type=int,38 nargs="+",39 default=[4096],40 help="Per-rank K values",41 )42 parser.add_argument("--num-workspaces", type=int, default=10)43 parser.add_argument("--warmup-replays", type=int, default=5)44 parser.add_argument("--samples", type=int, default=20)45 return parser.parse_args()46 47 48def make_gemm_rs(49 x: torch.Tensor,50 weight: torch.Tensor,51 partial: torch.Tensor,52 output: torch.Tensor,53 rows: int,54 device_group: dist.ProcessGroup,55) -> Callable[[], None]:56 def run() -> None:57 torch.mm(x, weight.T, out=partial[:rows])58 dist.reduce_scatter_single(output, partial, group=device_group)59 60 return run61 62 63def check_correctness(64 runs: dict[str, Callable[[], None]],65 outputs: dict[str, torch.Tensor],66 x: torch.Tensor,67 weight: torch.Tensor,68 padded_rows: int,69 rank: int,70 device_group: dist.ProcessGroup,71) -> None:72 expected_full = torch.zeros(73 (padded_rows, weight.shape[0]),74 dtype=x.dtype,75 device=x.device,76 )77 torch.mm(x, weight.T, out=expected_full[: x.shape[0]])78 dist.all_reduce(expected_full, group=device_group)79 expected = expected_full.chunk(dist.get_world_size(device_group))[rank]80 for name, run in runs.items():81 run()82 torch.accelerator.synchronize()83 torch.testing.assert_close(84 outputs[name],85 expected,86 rtol=5e-2,87 atol=4.0,88 )89 90 91def capture_graph(92 run: Callable[[], None],93 cpu_group: dist.ProcessGroup,94) -> tuple[torch.cuda.CUDAGraph, torch.cuda.Stream]:95 stream = torch.cuda.Stream()96 stream.wait_stream(torch.cuda.current_stream())97 dist.barrier(group=cpu_group)98 with torch.cuda.stream(stream):99 for _ in range(3):100 run()101 stream.synchronize()102 dist.barrier(group=cpu_group)103 104 graph = torch.cuda.CUDAGraph()105 with torch.cuda.graph(graph, stream=stream):106 run()107 torch.cuda.current_stream().wait_stream(stream)108 dist.barrier(group=cpu_group)109 return graph, stream110 111 112def benchmark_graphs(113 candidate_graphs: dict[str, list[torch.cuda.CUDAGraph]],114 warmup_replays: int,115 samples: int,116 device_group: dist.ProcessGroup,117 device_barrier: Callable[[], None],118) -> dict[str, float]:119 candidate_names = list(candidate_graphs)120 for round_index in range(warmup_replays):121 for candidate_index in range(len(candidate_names)):122 candidate_id = (round_index + candidate_index) % len(candidate_names)123 name = candidate_names[candidate_id]124 graphs = candidate_graphs[name]125 device_barrier()126 graphs[round_index % len(graphs)].replay()127 torch.accelerator.synchronize()128 129 timings: dict[str, list[float]] = {name: [] for name in candidate_names}130 start = torch.cuda.Event(enable_timing=True)131 end = torch.cuda.Event(enable_timing=True)132 for sample_index in range(samples):133 for candidate_index in range(len(candidate_names)):134 candidate_id = (sample_index + candidate_index) % len(candidate_names)135 name = candidate_names[candidate_id]136 graphs = candidate_graphs[name]137 device_barrier()138 start.record()139 graphs[sample_index % len(graphs)].replay()140 end.record()141 end.synchronize()142 143 elapsed_us = torch.tensor(144 start.elapsed_time(end) * 1000,145 dtype=torch.float64,146 device=torch.accelerator.current_device_index(),147 )148 dist.all_reduce(elapsed_us, op=dist.ReduceOp.MAX, group=device_group)149 timings[name].append(elapsed_us.item())150 return {name: statistics.median(values) for name, values in timings.items()}151 152 153def benchmark_shape(154 m: int,155 n: int,156 k: int,157 num_workspaces: int,158 warmup_replays: int,159 samples: int,160 device: torch.device,161 rank: int,162 world_size: int,163 device_group: dist.ProcessGroup,164 cpu_group: dist.ProcessGroup,165 device_barrier: Callable[[], None],166) -> dict[str, float | int]:167 padded_m = (m + world_size - 1) // world_size * world_size168 local_m = padded_m // world_size169 170 torch.manual_seed(1000 + rank * 10 + m + k)171 inputs = []172 weights = []173 for _ in range(num_workspaces):174 inputs.append(torch.randn(m, k, dtype=torch.bfloat16, device=device))175 weights.append(torch.randn(n, k, dtype=torch.bfloat16, device=device))176 177 ring_partial = torch.empty(padded_m, n, dtype=torch.bfloat16, device=device)178 ldmc_partial = symm_mem.empty(179 (padded_m, n),180 dtype=torch.bfloat16,181 device=device,182 )183 ldmc_handle = symm_mem.rendezvous(ldmc_partial, device_group)184 ring_output = torch.empty(local_m, n, dtype=torch.bfloat16, device=device)185 ldmc_output = torch.empty_like(ring_output)186 if padded_m > m:187 ring_partial[m:].zero_()188 ldmc_partial[m:].zero_()189 190 candidate_runs = {191 "ring_ll_us": [192 make_gemm_rs(193 x,194 weight,195 ring_partial,196 ring_output,197 m,198 device_group,199 )200 for x, weight in zip(inputs, weights)201 ],202 "ldmc_us": [203 make_gemm_rs(204 x,205 weight,206 ldmc_partial,207 ldmc_output,208 m,209 device_group,210 )211 for x, weight in zip(inputs, weights)212 ],213 }214 x = inputs[0]215 weight = weights[0]216 check_correctness(217 {name: runs[0] for name, runs in candidate_runs.items()},218 {"ring_ll_us": ring_output, "ldmc_us": ldmc_output},219 x,220 weight,221 padded_m,222 rank,223 device_group,224 )225 226 candidate_graphs = {}227 graph_keepalive: list[object] = [ldmc_handle]228 for name, runs in candidate_runs.items():229 bundles = [capture_graph(run, cpu_group) for run in runs]230 candidate_graphs[name] = [graph for graph, _ in bundles]231 graph_keepalive.extend(bundles)232 233 times = benchmark_graphs(234 candidate_graphs,235 warmup_replays,236 samples,237 device_group,238 device_barrier,239 )240 global_flops = 2 * m * n * k * world_size241 return {242 "M": m,243 "N": n,244 "K_per_rank": k,245 "K_global": k * world_size,246 **times,247 "ring_ll_tflops": global_flops / (times["ring_ll_us"] * 1e6),248 "ldmc_tflops": global_flops / (times["ldmc_us"] * 1e6),249 "ldmc_speedup": times["ring_ll_us"] / times["ldmc_us"],250 }251 252 253def main() -> None:254 args = parse_args()255 assert args.m and min(args.m) > 0256 assert args.k and min(args.k) > 0257 assert min(args.n, args.num_workspaces, args.samples) > 0258 assert args.warmup_replays >= 0259 local_rank = int(os.environ["LOCAL_RANK"])260 local_world_size = int(os.environ["LOCAL_WORLD_SIZE"])261 torch.accelerator.set_device_index(local_rank)262 init_distributed_environment()263 world_size = dist.get_world_size()264 os.environ["VLLM_ALLREDUCE_USE_SYMM_MEM"] = "0"265 symm_mem.set_backend("NCCL")266 with set_current_vllm_config(VllmConfig()):267 initialize_model_parallel(tensor_model_parallel_size=world_size)268 269 tp_group = get_tp_group()270 device_group = tp_group.device_group271 cpu_group = tp_group.cpu_group272 rank = tp_group.rank_in_group273 device = torch.device("cuda", local_rank)274 group_warmup = torch.zeros(1, device=device)275 dist.all_reduce(group_warmup, group=device_group)276 pynccl_comm = tp_group.device_communicator.pynccl_comm277 assert pynccl_comm is not None278 sync_input = torch.zeros(1, device=device)279 sync_output = torch.empty_like(sync_input)280 281 def device_barrier() -> None:282 # Order the timed launch after a device-side rank rendezvous without283 # including the rendezvous itself in the measured event interval.284 pynccl_comm.all_reduce(sync_input, sync_output)285 286 results = [287 benchmark_shape(288 m,289 args.n,290 k,291 args.num_workspaces,292 args.warmup_replays,293 args.samples,294 device,295 rank,296 world_size,297 device_group,298 cpu_group,299 device_barrier,300 )301 for k in args.k302 for m in args.m303 ]304 305 if rank == 0:306 metadata = {307 "world_size": world_size,308 "local_world_size": local_world_size,309 "num_nodes": world_size // local_world_size,310 "backend": dist.get_backend(device_group),311 "gpu": torch.cuda.get_device_name(local_rank),312 "torch": torch.__version__,313 "cuda": torch.version.cuda,314 }315 print(pd.Series(metadata, name="value").to_string())316 df = pd.DataFrame(results)317 print(df.to_string(index=False, float_format=lambda x: f"{x:.3f}"))318 319 dist.barrier(group=cpu_group)320 cleanup_dist_env_and_memory()321 322 323if __name__ == "__main__":324 main()325