From c076968c52cc8af147847c79352897ed0aef4e7e Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Sat, 21 Mar 2026 13:40:42 +0800 Subject: [PATCH] [CI] Remove obsolete AOT-only jit-kernel benchmarks after sgl-kernel 4.0 (#21075) --- .../benchmark/bench_awq_marlin_moe_repack.py | 125 --------- .../benchmark/bench_awq_marlin_repack.py | 110 -------- .../jit_kernel/benchmark/bench_gptq_marlin.py | 129 ---------- .../benchmark/bench_gptq_marlin_repack.py | 97 ------- .../benchmark/bench_moe_wna16_marlin.py | 240 ------------------ 5 files changed, 701 deletions(-) delete mode 100644 python/sglang/jit_kernel/benchmark/bench_awq_marlin_moe_repack.py delete mode 100644 python/sglang/jit_kernel/benchmark/bench_awq_marlin_repack.py delete mode 100644 python/sglang/jit_kernel/benchmark/bench_gptq_marlin.py delete mode 100644 python/sglang/jit_kernel/benchmark/bench_gptq_marlin_repack.py delete mode 100644 python/sglang/jit_kernel/benchmark/bench_moe_wna16_marlin.py diff --git a/python/sglang/jit_kernel/benchmark/bench_awq_marlin_moe_repack.py b/python/sglang/jit_kernel/benchmark/bench_awq_marlin_moe_repack.py deleted file mode 100644 index ac9faa517..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_awq_marlin_moe_repack.py +++ /dev/null @@ -1,125 +0,0 @@ -import numpy as np -import torch -import triton -import triton.testing -from sgl_kernel.scalar_type import scalar_types - -from sglang.jit_kernel.awq_marlin_repack import ( - awq_marlin_moe_repack as jit_awq_marlin_moe_repack, -) -from sglang.jit_kernel.benchmark.utils import run_benchmark -from sglang.srt.layers.quantization.utils import pack_cols, quantize_weights -from sglang.utils import is_in_ci - -AOT_AVAILABLE = hasattr(torch.ops.sgl_kernel, "awq_marlin_moe_repack") and hasattr( - torch.ops.sgl_kernel.awq_marlin_moe_repack, "default" -) - -IS_CI = is_in_ci() - -NUM_BITS = 4 -GROUP_SIZE = 128 -SIZE_N = 4096 - - -def awq_pack(q_w, num_bits, size_k, size_n): - if num_bits == 4: - interleave = np.array([0, 2, 4, 6, 1, 3, 5, 7]) - elif num_bits == 8: - interleave = np.array([0, 2, 1, 3]) - else: - raise Exception("num_bits must be 4 or 8, got {}".format(num_bits)) - - q_w = q_w.reshape((-1, len(interleave)))[:, interleave].ravel() - q_w = q_w.reshape((-1, size_n)).contiguous() - return pack_cols(q_w, num_bits, size_k, size_n) - - -def make_moe_weights(num_experts, size_k, size_n, num_bits, group_size): - pack_factor = 32 // num_bits - b_q_weight = torch.empty( - (num_experts, size_k, size_n // pack_factor), - dtype=torch.int32, - device="cuda", - ) - for e in range(num_experts): - b_weight = torch.randn((size_k, size_n), dtype=torch.float16, device="cuda") - w_ref, q_w, s, zp = quantize_weights( - b_weight, scalar_types.uint4, min(group_size, size_k), zero_points=True - ) - b_q_weight[e] = awq_pack(q_w, num_bits, size_k, size_n) - perm = torch.empty((num_experts, 0), dtype=torch.int32, device="cuda") - return b_q_weight, perm - - -def check_correctness(): - if not AOT_AVAILABLE: - print("sgl_kernel AOT not available, skipping correctness check") - return - - num_experts = 4 - size_k = 1024 - b_q_weight, perm = make_moe_weights( - num_experts, size_k, SIZE_N, NUM_BITS, GROUP_SIZE - ) - - out_jit = jit_awq_marlin_moe_repack(b_q_weight, perm, size_k, SIZE_N, NUM_BITS) - out_aot = torch.ops.sgl_kernel.awq_marlin_moe_repack.default( - b_q_weight, perm, size_k, SIZE_N, NUM_BITS - ) - torch.cuda.synchronize() - torch.testing.assert_close(out_jit, out_aot, rtol=0, atol=0) - print("Correctness check passed (JIT vs AOT)") - - -if IS_CI: - expert_range = [2, 4] -else: - expert_range = [2, 4, 8, 16] - -if AOT_AVAILABLE: - line_vals = ["jit", "aot"] - line_names = ["JIT Kernel", "AOT Kernel"] - styles = [("blue", "-"), ("green", "-")] -else: - line_vals = ["jit"] - line_names = ["JIT Kernel"] - styles = [("blue", "-")] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["num_experts"], - x_vals=expert_range, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="awq-marlin-moe-repack-performance", - args={"size_k": 4096, "size_n": SIZE_N, "num_bits": NUM_BITS}, - ) -) -def benchmark(num_experts, size_k, size_n, num_bits, provider): - group_size = min(GROUP_SIZE, size_k) - b_q_weight, perm = make_moe_weights( - num_experts, size_k, size_n, num_bits, group_size - ) - - if provider == "jit": - fn = lambda: jit_awq_marlin_moe_repack( - b_q_weight, perm, size_k, size_n, num_bits - ) - elif provider == "aot": - fn = lambda: torch.ops.sgl_kernel.awq_marlin_moe_repack.default( - b_q_weight, perm, size_k, size_n, num_bits - ) - else: - raise ValueError(f"Unknown provider: {provider}") - - return run_benchmark(fn) - - -if __name__ == "__main__": - check_correctness() - benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/benchmark/bench_awq_marlin_repack.py b/python/sglang/jit_kernel/benchmark/bench_awq_marlin_repack.py deleted file mode 100644 index 8cf748fc0..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_awq_marlin_repack.py +++ /dev/null @@ -1,110 +0,0 @@ -import numpy as np -import torch -import triton -import triton.testing -from sgl_kernel.scalar_type import scalar_types - -from sglang.jit_kernel.awq_marlin_repack import ( - awq_marlin_repack as jit_awq_marlin_repack, -) -from sglang.jit_kernel.benchmark.utils import run_benchmark -from sglang.srt.layers.quantization.utils import pack_cols, quantize_weights -from sglang.utils import is_in_ci - -AOT_AVAILABLE = hasattr(torch.ops.sgl_kernel, "awq_marlin_repack") and hasattr( - torch.ops.sgl_kernel.awq_marlin_repack, "default" -) - -IS_CI = is_in_ci() - -SIZE_K = 4096 -SIZE_N = 4096 -NUM_BITS = 4 -GROUP_SIZE = 128 - - -def awq_pack(q_w, num_bits, size_k, size_n): - if num_bits == 4: - interleave = np.array([0, 2, 4, 6, 1, 3, 5, 7]) - elif num_bits == 8: - interleave = np.array([0, 2, 1, 3]) - else: - raise Exception("num_bits must be 4 or 8, got {}".format(num_bits)) - - q_w = q_w.reshape((-1, len(interleave)))[:, interleave].ravel() - q_w = q_w.reshape((-1, size_n)).contiguous() - return pack_cols(q_w, num_bits, size_k, size_n) - - -_b_weight = torch.randn((SIZE_K, SIZE_N), dtype=torch.float16, device="cuda") -_w_ref, _q_w, _s, _zp = quantize_weights( - _b_weight, scalar_types.uint4, GROUP_SIZE, zero_points=True -) -_q_w_awq = awq_pack(_q_w, NUM_BITS, SIZE_K, SIZE_N) - - -def check_correctness(): - if not AOT_AVAILABLE: - print("sgl_kernel AOT not available, skipping correctness check") - return - out_jit = jit_awq_marlin_repack(_q_w_awq, SIZE_K, SIZE_N, NUM_BITS) - out_aot = torch.ops.sgl_kernel.awq_marlin_repack.default( - _q_w_awq, SIZE_K, SIZE_N, NUM_BITS - ) - torch.cuda.synchronize() - torch.testing.assert_close(out_jit, out_aot, rtol=0, atol=0) - print("Correctness check passed (JIT vs AOT)") - - -if IS_CI: - k_range = [1024, 4096] -else: - k_range = [512, 1024, 2048, 4096, 8192] - -if AOT_AVAILABLE: - line_vals = ["jit", "aot"] - line_names = ["JIT Kernel", "AOT Kernel"] - styles = [("blue", "-"), ("green", "-")] -else: - line_vals = ["jit"] - line_names = ["JIT Kernel"] - styles = [("blue", "-")] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["size_k"], - x_vals=k_range, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="awq-marlin-repack-performance", - args={"size_n": SIZE_N, "num_bits": NUM_BITS}, - ) -) -def benchmark(size_k, size_n, num_bits, provider): - group_size = min(GROUP_SIZE, size_k) - - b_weight = torch.randn((size_k, size_n), dtype=torch.float16, device="cuda") - w_ref, q_w, s, zp = quantize_weights( - b_weight, scalar_types.uint4, group_size, zero_points=True - ) - q_w_awq = awq_pack(q_w, num_bits, size_k, size_n) - - if provider == "jit": - fn = lambda: jit_awq_marlin_repack(q_w_awq, size_k, size_n, num_bits) - elif provider == "aot": - fn = lambda: torch.ops.sgl_kernel.awq_marlin_repack.default( - q_w_awq, size_k, size_n, num_bits - ) - else: - raise ValueError(f"Unknown provider: {provider}") - - return run_benchmark(fn) - - -if __name__ == "__main__": - check_correctness() - benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/benchmark/bench_gptq_marlin.py b/python/sglang/jit_kernel/benchmark/bench_gptq_marlin.py deleted file mode 100644 index 31881d785..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_gptq_marlin.py +++ /dev/null @@ -1,129 +0,0 @@ -import torch -import triton -import triton.testing -from sgl_kernel.scalar_type import scalar_types - -from sglang.jit_kernel.benchmark.utils import run_benchmark -from sglang.jit_kernel.gptq_marlin import gptq_marlin_gemm as jit_gptq_marlin_gemm -from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace -from sglang.test.test_marlin_utils import marlin_quantize -from sglang.utils import is_in_ci - -AOT_AVAILABLE = hasattr(torch.ops.sgl_kernel, "gptq_marlin_gemm") and hasattr( - torch.ops.sgl_kernel.gptq_marlin_gemm, "default" -) - -IS_CI = is_in_ci() - -SIZE_K = 4096 -SIZE_N = 4096 -GROUP_SIZE = 128 -QUANT_TYPE = scalar_types.uint4b8 - -_b_weight = torch.randn((SIZE_K, SIZE_N), dtype=torch.float16, device="cuda") -_w_ref, _marlin_q_w, _marlin_s, _g_idx, _sort_indices, _ = marlin_quantize( - _b_weight, QUANT_TYPE, GROUP_SIZE, act_order=False -) -_workspace = marlin_make_workspace(_w_ref.device) - - -def _run_gemm(fn, a): - return fn( - a, - None, - _marlin_q_w, - _marlin_s, - None, - None, - _g_idx, - _sort_indices, - _workspace, - QUANT_TYPE, - a.shape[0], - SIZE_N, - SIZE_K, - is_k_full=True, - use_atomic_add=False, - use_fp32_reduce=False, - is_zp_float=False, - ) - - -def _run_gemm_aot(a): - return torch.ops.sgl_kernel.gptq_marlin_gemm.default( - a, - None, - _marlin_q_w, - _marlin_s, - None, - None, - _g_idx, - _sort_indices, - _workspace, - QUANT_TYPE.id, - a.shape[0], - SIZE_N, - SIZE_K, - True, - False, - False, - False, - ) - - -def check_correctness(): - if not AOT_AVAILABLE: - print("sgl_kernel AOT not available, skipping correctness check") - return - a = torch.randn((16, SIZE_K), dtype=torch.float16, device="cuda") - out_jit = _run_gemm(jit_gptq_marlin_gemm, a) - out_aot = _run_gemm_aot(a) - torch.testing.assert_close(out_jit, out_aot, rtol=1e-3, atol=1e-3) - print("Correctness check passed (JIT vs AOT)") - - -if IS_CI: - m_range = [1, 16, 128] -else: - m_range = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512] - -if AOT_AVAILABLE: - line_vals = ["jit", "aot"] - line_names = ["JIT Kernel", "AOT Kernel"] - styles = [("blue", "-"), ("green", "-")] -else: - line_vals = ["jit"] - line_names = ["JIT Kernel"] - styles = [("blue", "-")] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["size_m"], - x_vals=m_range, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="gptq-marlin-gemm-performance", - args={}, - ) -) -def benchmark(size_m, provider): - device = torch.device("cuda") - a = torch.randn((size_m, SIZE_K), dtype=torch.float16, device=device) - - if provider == "jit": - fn = lambda: _run_gemm(jit_gptq_marlin_gemm, a) - elif provider == "aot": - fn = lambda: _run_gemm_aot(a) - else: - raise ValueError(f"Unknown provider: {provider}") - - return run_benchmark(fn) - - -if __name__ == "__main__": - check_correctness() - benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/benchmark/bench_gptq_marlin_repack.py b/python/sglang/jit_kernel/benchmark/bench_gptq_marlin_repack.py deleted file mode 100644 index 3fe86a0e6..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_gptq_marlin_repack.py +++ /dev/null @@ -1,97 +0,0 @@ -import torch -import triton -import triton.testing -from sgl_kernel.scalar_type import scalar_types - -from sglang.jit_kernel.benchmark.utils import run_benchmark -from sglang.jit_kernel.gptq_marlin_repack import gptq_marlin_repack as jit_fn -from sglang.srt.layers.quantization.utils import gptq_quantize_weights, pack_rows -from sglang.utils import is_in_ci - -AOT_AVAILABLE = hasattr(torch.ops.sgl_kernel, "gptq_marlin_repack") and hasattr( - torch.ops.sgl_kernel.gptq_marlin_repack, "default" -) - -IS_CI = is_in_ci() - -SIZE_N = 4096 -NUM_BITS = 4 -QUANT_TYPE = scalar_types.uint4b8 -GROUP_SIZE = 128 - -_cache = {} - - -def _get_inputs(size_k): - if size_k not in _cache: - size_n = SIZE_N - b_weight = torch.randn((size_k, size_n), dtype=torch.float16, device="cuda") - _, q_w, _, _, _ = gptq_quantize_weights( - b_weight, QUANT_TYPE, GROUP_SIZE, act_order=False - ) - q_w_gptq = pack_rows(q_w, NUM_BITS, size_k, size_n) - sort_indices = torch.empty(0, dtype=torch.int, device="cuda") - _cache[size_k] = (q_w_gptq, sort_indices) - return _cache[size_k] - - -def check_correctness(): - if not AOT_AVAILABLE: - print("sgl_kernel AOT not available, skipping correctness check") - return - size_k = 4096 - q_w_gptq, sort_indices = _get_inputs(size_k) - out_jit = jit_fn(q_w_gptq, sort_indices, size_k, SIZE_N, NUM_BITS) - out_aot = torch.ops.sgl_kernel.gptq_marlin_repack.default( - q_w_gptq, sort_indices, size_k, SIZE_N, NUM_BITS - ) - torch.testing.assert_close(out_jit, out_aot, rtol=0, atol=0) - print("Correctness check passed (JIT vs AOT)") - - -if IS_CI: - k_range = [128, 1024, 4096] -else: - k_range = [128, 256, 512, 1024, 2048, 4096, 8192] - -if AOT_AVAILABLE: - line_vals = ["jit", "aot"] - line_names = ["JIT Kernel", "AOT Kernel"] - styles = [("blue", "-"), ("green", "-")] -else: - line_vals = ["jit"] - line_names = ["JIT Kernel"] - styles = [("blue", "-")] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["size_k"], - x_vals=k_range, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="gptq-marlin-repack-performance", - args={}, - ) -) -def benchmark(size_k, provider): - q_w_gptq, sort_indices = _get_inputs(size_k) - - if provider == "jit": - fn = lambda: jit_fn(q_w_gptq, sort_indices, size_k, SIZE_N, NUM_BITS) - elif provider == "aot": - fn = lambda: torch.ops.sgl_kernel.gptq_marlin_repack.default( - q_w_gptq, sort_indices, size_k, SIZE_N, NUM_BITS - ) - else: - raise ValueError(f"Unknown provider: {provider}") - - return run_benchmark(fn) - - -if __name__ == "__main__": - check_correctness() - benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/benchmark/bench_moe_wna16_marlin.py b/python/sglang/jit_kernel/benchmark/bench_moe_wna16_marlin.py deleted file mode 100644 index fa6e9b36d..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_moe_wna16_marlin.py +++ /dev/null @@ -1,240 +0,0 @@ -import torch -import triton -import triton.testing -from sgl_kernel.scalar_type import scalar_types - -from sglang.jit_kernel.benchmark.utils import run_benchmark -from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm as jit_fn -from sglang.srt.layers.moe.fused_moe_triton import moe_align_block_size -from sglang.test.test_marlin_utils import marlin_quantize -from sglang.utils import is_in_ci - -AOT_AVAILABLE = hasattr(torch.ops.sgl_kernel, "moe_wna16_marlin_gemm") and hasattr( - torch.ops.sgl_kernel.moe_wna16_marlin_gemm, "default" -) - -IS_CI = is_in_ci() - - -def stack_and_dev(tensors): - dev = tensors[0].device - return torch.stack(tensors, dim=0).to(dev) - - -E = 8 -SIZE_K = 4096 -SIZE_N = 4096 -GROUP_SIZE = 128 -TOPK = 2 -QUANT_TYPE = scalar_types.uint4b8 -DTYPE = torch.float16 -BLOCK_SIZE_M = 64 - -torch.manual_seed(0) -_qweight_l, _scales_l, _w_ref_l = [], [], [] -for i in range(E): - _w = torch.randn((SIZE_N, SIZE_K), dtype=DTYPE, device="cuda") / 20 - _perm = torch.randperm(SIZE_K) - _w_ref, _qw, _s, _, _, _ = marlin_quantize(_w, QUANT_TYPE, GROUP_SIZE, False, _perm) - _w_ref_l.append(_w_ref.T) - _qweight_l.append(_qw) - _scales_l.append(_s) - -_qweight = stack_and_dev(_qweight_l).contiguous() -_scales = stack_and_dev(_scales_l) - -_sms = torch.cuda.get_device_properties("cuda").multi_processor_count - - -def _make_inputs(size_m): - a = torch.randn((size_m, SIZE_K), dtype=DTYPE, device="cuda") / 10 - score = torch.randn((size_m, E), dtype=DTYPE, device="cuda") - score_softmax = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weights, topk_ids = torch.topk(score_softmax, TOPK) - - sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size( - topk_ids, BLOCK_SIZE_M, E - ) - - max_workspace_size = (SIZE_N // 64) * (sorted_token_ids.size(0) // BLOCK_SIZE_M) - max_workspace_size = min(max_workspace_size, _sms * 4) - workspace = torch.zeros(max_workspace_size, dtype=torch.int, device="cuda") - - c = torch.empty((size_m * TOPK, SIZE_N), dtype=DTYPE, device="cuda") - - return ( - a, - c, - topk_weights, - topk_ids, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - workspace, - ) - - -def _run_jit( - a, - c, - topk_weights, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - workspace, - size_m, -): - return jit_fn( - a, - c, - _qweight, - None, - _scales, - None, - None, - None, - None, - workspace, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - topk_weights, - moe_block_size=BLOCK_SIZE_M, - top_k=TOPK, - mul_topk_weights=False, - is_ep=False, - b_q_type=QUANT_TYPE, - size_m=size_m, - size_n=SIZE_N, - size_k=SIZE_K, - is_k_full=True, - use_atomic_add=True, - use_fp32_reduce=True, - is_zp_float=False, - ) - - -def _run_aot( - a, - c, - topk_weights, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - workspace, - size_m, -): - return torch.ops.sgl_kernel.moe_wna16_marlin_gemm.default( - a, - c, - _qweight, - None, - _scales, - None, - None, - None, - None, - workspace, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - topk_weights, - moe_block_size=BLOCK_SIZE_M, - top_k=TOPK, - mul_topk_weights=False, - is_ep=False, - b_q_type_id=QUANT_TYPE.id, - size_m=size_m, - size_n=SIZE_N, - size_k=SIZE_K, - is_k_full=True, - use_atomic_add=True, - use_fp32_reduce=True, - is_zp_float=False, - ) - - -def check_correctness(): - if not AOT_AVAILABLE: - print("sgl_kernel AOT not available, skipping correctness check") - return - size_m = 16 - a, c, topk_weights, topk_ids, sorted_token_ids, expert_ids, ntp, workspace = ( - _make_inputs(size_m) - ) - c_jit = c.clone() - c_aot = c.clone() - _run_jit( - a, c_jit, topk_weights, sorted_token_ids, expert_ids, ntp, workspace, size_m - ) - _run_aot( - a, c_aot, topk_weights, sorted_token_ids, expert_ids, ntp, workspace, size_m - ) - torch.testing.assert_close(c_jit, c_aot, rtol=1e-3, atol=1e-3) - print("Correctness check passed (JIT vs AOT)") - - -if IS_CI: - m_range = [1, 16, 128] -else: - m_range = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512] - -if AOT_AVAILABLE: - line_vals = ["jit", "aot"] - line_names = ["JIT Kernel", "AOT Kernel"] - styles = [("blue", "-"), ("green", "-")] -else: - line_vals = ["jit"] - line_names = ["JIT Kernel"] - styles = [("blue", "-")] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["size_m"], - x_vals=m_range, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="moe-wna16-marlin-gemm-performance", - args={}, - ) -) -def benchmark(size_m, provider): - a, c, topk_weights, topk_ids, sorted_token_ids, expert_ids, ntp, workspace = ( - _make_inputs(size_m) - ) - - if provider == "jit": - fn = lambda: _run_jit( - a, - c.clone(), - topk_weights, - sorted_token_ids, - expert_ids, - ntp, - workspace, - size_m, - ) - elif provider == "aot": - fn = lambda: _run_aot( - a, - c.clone(), - topk_weights, - sorted_token_ids, - expert_ids, - ntp, - workspace, - size_m, - ) - else: - raise ValueError(f"Unknown provider: {provider}") - - return run_benchmark(fn) - - -if __name__ == "__main__": - check_correctness() - benchmark.run(print_data=True)