[CI] Remove obsolete AOT-only jit-kernel benchmarks after sgl-kernel 4.0 (#21075)
This commit is contained in:
@@ -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)
|
|
||||||
@@ -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)
|
|
||||||
@@ -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)
|
|
||||||
@@ -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)
|
|
||||||
@@ -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)
|
|
||||||
Reference in New Issue
Block a user