diff --git a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py index f45a230ee..1fa3819cc 100644 --- a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py +++ b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py @@ -337,13 +337,13 @@ def parse_args(): parser.add_argument( "--prefill-shapes", type=str, - default="2048x8192,8192x8192,16384x8192", + default="2048x2880,2048x8192,8192x8192,16384x8192", help="Comma-separated MxN shapes for eager mode.", ) parser.add_argument( "--decode-shapes", type=str, - default="1x8192,2x8192,4x8192,8x8192,16x8192", + default="1x2880,4x2880,16x2880,1x8192,2x8192,4x8192,8x8192,16x8192", help="Comma-separated MxN shapes for graph mode.", ) parser.add_argument("--warmup", type=int, default=10) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 98be02687..dad005fc2 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -628,7 +628,7 @@ class GroupCoordinator: weight_: torch.Tensor, eps: float, ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: - """Attempt fused all-reduce + RMSNorm via custom all-reduce communicator.""" + """Attempt fused all-reduce + RMSNorm via custom all-reduce communicator. ROCm/HIP Only""" ca_comm = self.ca_comm if ca_comm is None or getattr(ca_comm, "disabled", True): return None @@ -646,24 +646,17 @@ class GroupCoordinator: if not hasattr(ca_comm, "custom_fused_ar_rms"): return None - # 1-stage policy for fused AR+RMSNorm: - # 1) Explicit env override wins. - # 2) Deterministic inference forces 1-stage for reproducibility. - # 3) Otherwise follow AITER's heuristic (small payloads only). + # 1-stage vs 2-stage selection for fused AR+RMSNorm: + # The 1-stage kernel launches one block per token and is capped at + # 80 tokens (kMaxBlocks). Guard with a byte threshold so large + # prefill batches fall through to the 2-stage kernel instead of + # hitting a runtime error. AITER's C++ dispatch already gates + # which hidden_dims have valid 1-stage support. if envs.SGLANG_USE_1STAGE_ALLREDUCE.is_set(): use_1stage_ar = envs.SGLANG_USE_1STAGE_ALLREDUCE.get() - elif envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get(): - use_1stage_ar = True else: total_bytes = input_.numel() * input_.element_size() - hidden_dim = input_.shape[-1] - use_1stage_ar = total_bytes <= 128 * 1024 and hidden_dim in { - 512, - 1024, - 2048, - 2880, - 4096, - } + use_1stage_ar = total_bytes <= 128 * 1024 fused_outputs = ca_comm.custom_fused_ar_rms( input_, diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 05308af3e..54e2326b2 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -167,11 +167,12 @@ def apply_flashinfer_allreduce_fusion(batch_size: int): def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor): n = input_tensor.shape[-1] total_bytes = input_tensor.numel() * input_tensor.element_size() + # Aiter's should_custom_ar uses <= max_size/2 (64 MB); match that boundary. return ( _use_aiter and total_bytes > 0 and n <= 16384 - and total_bytes < 8 * 1024 * 8192 + and total_bytes <= 8 * 1024 * 8192 and get_tensor_model_parallel_world_size() != 6 and not is_dp_attention_enabled() and get_global_server_args().enable_aiter_allreduce_fusion diff --git a/test/registered/ops/test_aiter_allreduce_fusion_amd.py b/test/registered/ops/test_aiter_allreduce_fusion_amd.py index 3fe3e9b19..cf1b201fa 100644 --- a/test/registered/ops/test_aiter_allreduce_fusion_amd.py +++ b/test/registered/ops/test_aiter_allreduce_fusion_amd.py @@ -12,14 +12,142 @@ from sglang.test.ci.ci_register import register_amd_ci register_amd_ci(est_time=240, suite="stage-c-test-large-8-gpu-amd") +HIDDEN_DIMS = [2880, 4096, 5120, 6144, 7168, 8192] + + +def _run_residual_accuracy_check(): + """Distributed entry point: bit-exact residual accuracy across 1-stage/2-stage. + + Regression test for the 1-stage kernel accuracy bug (ROCm/aiter#2586): + allreduce_fusion_kernel_1stage accumulated in f32 and added the residual + before rounding to bf16, while the unfused path rounds allreduce to bf16 + first. The 1-ULP divergence compounded across layers and caused a -2.6pp + GSM8K regression. + + Must be launched via torchrun (multi-GPU). + """ + import torch.distributed as dist + + from sglang.srt.distributed.communication_op import ( + tensor_model_parallel_all_reduce, + tensor_model_parallel_fused_allreduce_rmsnorm, + ) + from sglang.srt.distributed.parallel_state import ( + destroy_distributed_environment, + destroy_model_parallel, + init_distributed_environment, + initialize_model_parallel, + set_custom_all_reduce, + ) + + rank = int(os.environ.get("RANK", "0")) + world_size = int(os.environ.get("WORLD_SIZE", "1")) + local_rank = int(os.environ.get("LOCAL_RANK", str(rank))) + torch.cuda.set_device(local_rank % torch.cuda.device_count()) + device = torch.device(f"cuda:{local_rank % torch.cuda.device_count()}") + + set_custom_all_reduce(True) + init_distributed_environment( + world_size=world_size, + rank=rank, + local_rank=local_rank, + distributed_init_method="env://", + backend="nccl", + ) + initialize_model_parallel(tensor_model_parallel_size=world_size) + + dtype = torch.bfloat16 + eps = 1e-6 + + all_pass = True + test_cases = [(m, n) for n in HIDDEN_DIMS for m in [1, 4, 8, 16, 32, 64, 128]] + + prev_n = None + for m, n in test_cases: + if n != prev_n: + prev_n = n + weight = torch.ones((n,), dtype=dtype, device=device) + if rank == 0: + print(f"\nhidden_dim={n}:") + + torch.manual_seed(1234 + rank * 17 + m) + x = torch.randn((m, n), dtype=torch.float32, device=device).to(dtype) + residual = torch.randn((m, n), dtype=torch.float32, device=device).to(dtype) + zero_res = torch.zeros((m, n), dtype=dtype, device=device) + + dist.barrier() + torch.cuda.synchronize() + + fused_zero = tensor_model_parallel_fused_allreduce_rmsnorm( + x.clone(), zero_res.clone(), weight, eps + ) + torch.cuda.synchronize() + if fused_zero is None: + if rank == 0: + print(f" {m:>5d}x{n}: SKIP (fused unavailable)") + continue + _, fused_ar = fused_zero + + dist.barrier() + torch.cuda.synchronize() + + fused_random = tensor_model_parallel_fused_allreduce_rmsnorm( + x.clone(), residual.clone(), weight, eps + ) + torch.cuda.synchronize() + _, fused_res = fused_random + + dist.barrier() + torch.cuda.synchronize() + + unfused_ar = tensor_model_parallel_all_reduce(x.clone()) + torch.cuda.synchronize() + + expected = fused_ar + residual + diff = (fused_res.float() - expected.float()).abs() + ar_diff = (fused_ar.float() - unfused_ar.float()).abs() + max_diff = diff.max().item() + frac_nonzero = (diff > 0).float().mean().item() + + nbytes = m * n * dtype.itemsize + stage = "1-stage" if nbytes <= 128 * 1024 else "2-stage" + passed = max_diff == 0.0 + + if not passed: + all_pass = False + + if rank == 0: + status = "PASS" if passed else "FAIL" + print( + f" {m:>5d}x{n} ({stage:>7s}): max_diff={max_diff:.6e} " + f"frac_nonzero={frac_nonzero:.4f} " + f"AR_exact={'yes' if ar_diff.max().item() == 0 else 'no':>3s} " + f"[{status}]" + ) + + dist.barrier() + destroy_model_parallel() + destroy_distributed_environment() + + if rank == 0: + print() + if all_pass: + print("ALL PASSED: fused residual output is bit-identical to unfused path.") + else: + print( + "FAILED: fused residual output diverges from unfused path for some shapes." + ) + sys.exit(0 if all_pass else 1) + class TestAiterAllreduceFusionAmd(unittest.TestCase): - def test_fused_ar_rms_benchmark(self): - if not torch.cuda.is_available(): - self.skipTest("CUDA/ROCm device is not available.") - if torch.cuda.device_count() < 8: - self.skipTest("This test requires at least 8 GPUs.") + @staticmethod + def _gpu_count(): + return torch.cuda.device_count() if torch.cuda.is_available() else 0 + + def _run_benchmark(self, nproc, prefill_shapes, decode_shapes): + """Run the benchmark subprocess and return parsed CSV rows.""" repo_root = Path(__file__).resolve().parents[3] benchmark_script = ( repo_root @@ -40,15 +168,14 @@ class TestAiterAllreduceFusionAmd(unittest.TestCase): "-m", "torch.distributed.run", "--standalone", - "--nproc_per_node=8", + f"--nproc_per_node={nproc}", str(benchmark_script), "--dtype", "bf16", "--prefill-shapes", - # Include both <=64MiB and >64MiB shapes to verify default gate behavior. - "128x7168,512x7168,2048x7168,4096x7168,5120x7168", + prefill_shapes, "--decode-shapes", - "1x7168,8x7168,64x7168,512x7168", + decode_shapes, "--warmup", "3", "--iters", @@ -84,39 +211,125 @@ class TestAiterAllreduceFusionAmd(unittest.TestCase): rows = list(csv.DictReader(f)) self.assertGreater(len(rows), 0, "CSV contains no rows.") + return rows - eager_rows = [r for r in rows if r["mode"] == "eager"] - graph_rows = [r for r in rows if r["mode"] == "graph"] - self.assertGreater(len(eager_rows), 0, "Missing eager rows in CSV.") - self.assertGreater(len(graph_rows), 0, "Missing graph rows in CSV.") + def _assert_correctness(self, rows): + bad_rows = [r for r in rows if r["correctness_ok"] != "True"] + self.assertEqual( + [], + bad_rows, + f"Found correctness failures: {bad_rows}", + ) - # Correctness should always pass regardless of fused availability. - bad_rows = [r for r in rows if r["correctness_ok"] != "True"] - self.assertEqual( - [], - bad_rows, - f"Found correctness failures: {bad_rows}", + def test_fused_ar_rms_benchmark(self): + if self._gpu_count() < 8: + self.skipTest("This test requires at least 8 GPUs.") + + rows = self._run_benchmark( + nproc=8, + prefill_shapes="128x7168,512x7168,2048x7168,4096x7168,5120x7168", + decode_shapes="1x7168,8x7168,64x7168,512x7168", + ) + + eager_rows = [r for r in rows if r["mode"] == "eager"] + graph_rows = [r for r in rows if r["mode"] == "graph"] + self.assertGreater(len(eager_rows), 0, "Missing eager rows in CSV.") + self.assertGreater(len(graph_rows), 0, "Missing graph rows in CSV.") + + self._assert_correctness(rows) + + self.assertTrue( + any(r["fused_available"] == "True" for r in eager_rows), + "Expected at least one eager row with fused_available=True.", + ) + self.assertTrue( + any(r["fused_available"] == "True" for r in graph_rows), + "Expected at least one graph row with fused_available=True.", + ) + + large_eager_rows = [ + r for r in eager_rows if int(r["bytes_per_rank"]) > 64 * 1024 * 1024 + ] + self.assertTrue( + any(r["fused_available"] == "False" for r in large_eager_rows), + "Expected fused fallback for oversized eager shape(s) under default gate.", + ) + + def test_fused_ar_rms_multi_hidden_dim(self): + """Correctness across hidden_dims from various models (TP=4).""" + nproc = min(self._gpu_count(), 4) + if nproc < 2: + self.skipTest("This test requires at least 2 GPUs.") + + # hidden_dims: 2880 (GPT-OSS), 4096 (Qwen3.5), 5120, 6144 (Mixtral), + # 7168 (DeepSeek), 8192 (Llama-70B) + decode = ",".join(f"{m}x{n}" for n in HIDDEN_DIMS for m in [1, 4, 16]) + prefill = ",".join(f"128x{n}" for n in HIDDEN_DIMS) + + rows = self._run_benchmark( + nproc=nproc, + prefill_shapes=prefill, + decode_shapes=decode, + ) + + self._assert_correctness(rows) + + fused_rows = [r for r in rows if r["fused_available"] == "True"] + self.assertEqual( + len(fused_rows), + len(rows), + f"Expected fused available for all shapes, but {len(rows) - len(fused_rows)} " + f"rows were not fused: " + f"{[r['shape'] for r in rows if r['fused_available'] != 'True']}", + ) + + def test_fused_ar_rms_residual_accuracy(self): + """Bit-exact residual accuracy across 1-stage and 2-stage paths. + + Regression test for ROCm/aiter#2586. Launches this file itself via + torchrun with --residual-accuracy to run the distributed check. + """ + nproc = min(self._gpu_count(), 4) + if nproc < 2: + self.skipTest("This test requires at least 2 GPUs.") + + cmd = [ + sys.executable, + "-m", + "torch.distributed.run", + "--standalone", + f"--nproc_per_node={nproc}", + __file__, + "--residual-accuracy", + ] + + result = subprocess.run( + cmd, + cwd=str(Path(__file__).resolve().parents[3]), + env=os.environ.copy(), + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + timeout=300, + ) + + if result.returncode != 0: + self.fail( + "Residual accuracy check failed.\n" + f"Return code: {result.returncode}\n" + f"Command: {' '.join(cmd)}\n" + f"Output:\n{result.stdout}" ) - # We should see fused path active for small shapes in both modes. - self.assertTrue( - any(r["fused_available"] == "True" for r in eager_rows), - "Expected at least one eager row with fused_available=True.", - ) - self.assertTrue( - any(r["fused_available"] == "True" for r in graph_rows), - "Expected at least one graph row with fused_available=True.", - ) - - # Default gate should reject at least one oversized eager shape. - large_eager_rows = [ - r for r in eager_rows if int(r["bytes_per_rank"]) > 64 * 1024 * 1024 - ] - self.assertTrue( - any(r["fused_available"] == "False" for r in large_eager_rows), - "Expected fused fallback for oversized eager shape(s) under default gate.", - ) + self.assertIn( + "ALL PASSED", + result.stdout, + f"Expected 'ALL PASSED' in output, got:\n{result.stdout}", + ) if __name__ == "__main__": - unittest.main() + if "--residual-accuracy" in sys.argv: + _run_residual_accuracy_check() + else: + unittest.main()