[Kernel] Remove unused implementations and stale registry entries (#32636)
This commit is contained in:
@@ -1,292 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
|
||||
from sglang.kernels.ops.quantization.mxfp8 import (
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant,
|
||||
es_sm100_mxfp8_blockscaled_moe_grouped_gemm,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(
|
||||
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
|
||||
)
|
||||
|
||||
|
||||
def is_sm100_supported(device=None) -> bool:
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
return (torch.cuda.get_device_capability(device)[0] == 10) and (
|
||||
torch.version.cuda >= "12.8"
|
||||
)
|
||||
|
||||
|
||||
_SM100_SUPPORTED = is_sm100_supported()
|
||||
|
||||
|
||||
def _probe_sgl_kernel_group_mm() -> tuple[bool, str]:
|
||||
if not _SM100_SUPPORTED:
|
||||
return False, "MXFP8 MoE benchmark requires sm100+ with CUDA 12.8+."
|
||||
try:
|
||||
import sgl_kernel # noqa: F401
|
||||
except Exception as e:
|
||||
return False, f"import sgl_kernel failed: {e}"
|
||||
if not hasattr(sgl_kernel, "es_sm100_mxfp8_blockscaled_grouped_mm"):
|
||||
return False, "sgl_kernel.es_sm100_mxfp8_blockscaled_grouped_mm is missing."
|
||||
try:
|
||||
pass
|
||||
|
||||
# We assume if it's imported, it works
|
||||
except Exception as e:
|
||||
return False, f"calling sgl-kernel grouped_mm op failed: {e}"
|
||||
return True, ""
|
||||
|
||||
|
||||
_SGL_KERNEL_AVAILABLE, _SGL_KERNEL_REASON = _probe_sgl_kernel_group_mm()
|
||||
|
||||
|
||||
def align(val: int, alignment: int = 128) -> int:
|
||||
return int((val + alignment - 1) // alignment * alignment)
|
||||
|
||||
|
||||
def _prepare_case(
|
||||
total_tokens: int, n_g: int, k_g: int, num_experts: int, dtype: torch.dtype
|
||||
) -> dict[str, Any]:
|
||||
device = torch.device("cuda")
|
||||
base = total_tokens // num_experts
|
||||
rem = total_tokens % num_experts
|
||||
m_per_expert = [base + (1 if i < rem else 0) for i in range(num_experts)]
|
||||
|
||||
expert_offset = 0
|
||||
expert_offsets = []
|
||||
aux_expert_offset = 0
|
||||
aux_expert_offsets = []
|
||||
a_blockscale_offset = 0
|
||||
a_blockscale_offsets = []
|
||||
b_blockscale_offset = 0
|
||||
b_blockscale_offsets = []
|
||||
tokens_per_expert_list = []
|
||||
expert_ranges = []
|
||||
problem_sizes = []
|
||||
|
||||
a_list = []
|
||||
b_list = []
|
||||
for g in range(num_experts):
|
||||
m_g = m_per_expert[g]
|
||||
tokens_per_expert_list.append(m_g)
|
||||
expert_ranges.append((expert_offset, expert_offset + m_g))
|
||||
expert_offsets.append(expert_offset)
|
||||
expert_offset += m_g
|
||||
|
||||
aux_expert_offsets.append(aux_expert_offset)
|
||||
aux_expert_offset += n_g
|
||||
|
||||
a_blockscale_offsets.append(a_blockscale_offset)
|
||||
a_blockscale_offset += align(m_g, 128)
|
||||
|
||||
b_blockscale_offsets.append(b_blockscale_offset)
|
||||
b_blockscale_offset += n_g # n_g already align to 128 in practice
|
||||
|
||||
problem_sizes.append([m_g, n_g, k_g])
|
||||
|
||||
a = torch.randn((m_g, k_g), device=device, dtype=dtype) * 0.1
|
||||
b = torch.randn((n_g, k_g), device=device, dtype=dtype) * 0.1
|
||||
a_list.append(a)
|
||||
b_list.append(b)
|
||||
|
||||
a = torch.concat(a_list, dim=0)
|
||||
b = torch.concat(b_list, dim=0)
|
||||
|
||||
_expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32)
|
||||
_aux_expert_offsets = torch.tensor(aux_expert_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_tokens_per_expert = torch.tensor(tokens_per_expert_list).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_problem_sizes = torch.tensor(problem_sizes).to(device=device, dtype=torch.int32)
|
||||
|
||||
a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device)
|
||||
a_scale_factor = torch.zeros(
|
||||
(a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device
|
||||
)
|
||||
|
||||
b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device)
|
||||
b_scale_factor = torch.zeros(
|
||||
(num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device
|
||||
)
|
||||
|
||||
# Use a global workspace to avoid allocating 1GB every time
|
||||
workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device)
|
||||
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant(
|
||||
a,
|
||||
_tokens_per_expert,
|
||||
_expert_offsets,
|
||||
_a_blockscale_offsets,
|
||||
a_quant,
|
||||
a_scale_factor,
|
||||
)
|
||||
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant(
|
||||
b,
|
||||
torch.ones_like(_tokens_per_expert) * n_g,
|
||||
_aux_expert_offsets,
|
||||
_b_blockscale_offsets,
|
||||
b_quant,
|
||||
b_scale_factor,
|
||||
)
|
||||
|
||||
b_quant = b_quant.view(num_experts, n_g, k_g)
|
||||
b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32)
|
||||
|
||||
sgl_b_quant = b_quant.transpose(1, 2)
|
||||
sgl_b_scale_factor = b_scale_factor.transpose(1, 2)
|
||||
|
||||
return {
|
||||
"a": a,
|
||||
"b": b.view(num_experts, n_g, k_g),
|
||||
"b_quant": b_quant,
|
||||
"a_quant": a_quant,
|
||||
"b_scale_factor": b_scale_factor,
|
||||
"a_scale_factor": a_scale_factor,
|
||||
"expert_offsets": _expert_offsets,
|
||||
"a_blockscale_offsets": _a_blockscale_offsets,
|
||||
"tokens_per_expert": _tokens_per_expert,
|
||||
"problem_sizes": _problem_sizes,
|
||||
"sgl_b_quant": sgl_b_quant,
|
||||
"sgl_b_scale_factor": sgl_b_scale_factor,
|
||||
"workspace": workspace,
|
||||
"expert_ranges": expert_ranges,
|
||||
"dtype": dtype,
|
||||
}
|
||||
|
||||
|
||||
def _sgl_kernel_group_mm(case: dict[str, Any]) -> torch.Tensor:
|
||||
from sgl_kernel import es_sm100_mxfp8_blockscaled_grouped_mm
|
||||
|
||||
a_quant = case["a_quant"]
|
||||
sgl_b_quant = case["sgl_b_quant"]
|
||||
a_scale_factor = case["a_scale_factor"]
|
||||
sgl_b_scale_factor = case["sgl_b_scale_factor"]
|
||||
problem_sizes = case["problem_sizes"]
|
||||
expert_offsets = case["expert_offsets"]
|
||||
a_blockscale_offsets = case["a_blockscale_offsets"]
|
||||
dtype = case["dtype"]
|
||||
|
||||
total_tokens = a_quant.shape[0]
|
||||
n_g = sgl_b_quant.shape[2]
|
||||
|
||||
# sgl-kernel takes output pre-allocated
|
||||
d = torch.empty((total_tokens, n_g), device=a_quant.device, dtype=dtype)
|
||||
es_sm100_mxfp8_blockscaled_grouped_mm(
|
||||
d,
|
||||
a_quant,
|
||||
sgl_b_quant,
|
||||
a_scale_factor,
|
||||
sgl_b_scale_factor,
|
||||
problem_sizes,
|
||||
expert_offsets,
|
||||
a_blockscale_offsets,
|
||||
)
|
||||
return d
|
||||
|
||||
|
||||
shape_range = get_benchmark_range(
|
||||
full_range=[
|
||||
# (total_tokens, n_g, k_g, num_experts)
|
||||
(1024, 4096, 4096, 64),
|
||||
(2048, 4096, 4096, 64),
|
||||
(4096, 4096, 4096, 64),
|
||||
]
|
||||
+ [
|
||||
(total_tokens, n_g, k_g, num_experts)
|
||||
for total_tokens in [32 * (2**i) for i in range(9)] # 32 to 8192
|
||||
for n_g, k_g, num_experts in [
|
||||
# DeepSeek-V3/R1, gateup, TP = 1, EP = 8
|
||||
(4096, 7168, 32),
|
||||
# DeepSeek-V3/R1, down, TP = 1, EP = 8
|
||||
(7168, 2048, 32),
|
||||
]
|
||||
],
|
||||
ci_range=[(1024, 2048, 2048, 8)],
|
||||
)
|
||||
|
||||
line_vals = ["jit"]
|
||||
line_names = ["JIT MXFP8 MoE GroupMM"]
|
||||
styles = [("green", "-")]
|
||||
|
||||
if _SGL_KERNEL_AVAILABLE:
|
||||
line_vals.append("sgl_kernel")
|
||||
line_names.append("sgl-kernel MXFP8 MoE GroupMM")
|
||||
styles.append(("orange", "-"))
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["total_tokens", "n_g", "k_g", "num_experts"],
|
||||
x_vals=shape_range,
|
||||
x_log=False,
|
||||
line_arg="provider",
|
||||
line_vals=line_vals,
|
||||
line_names=line_names,
|
||||
styles=styles,
|
||||
ylabel="us",
|
||||
plot_name="mxfp8-moe-groupmm-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(total_tokens, n_g, k_g, num_experts, provider):
|
||||
case = _prepare_case(total_tokens, n_g, k_g, num_experts, torch.bfloat16)
|
||||
|
||||
if provider == "jit":
|
||||
fn = lambda: es_sm100_mxfp8_blockscaled_moe_grouped_gemm(
|
||||
case["b_quant"],
|
||||
case["a_quant"],
|
||||
case["b_scale_factor"],
|
||||
case["a_scale_factor"],
|
||||
case["expert_offsets"],
|
||||
case["a_blockscale_offsets"],
|
||||
case["tokens_per_expert"],
|
||||
case["workspace"],
|
||||
case["dtype"],
|
||||
)
|
||||
elif provider == "sgl_kernel":
|
||||
fn = lambda: _sgl_kernel_group_mm(case)
|
||||
else:
|
||||
raise ValueError(f"Unknown provider: {provider}")
|
||||
|
||||
# Warm up
|
||||
fn()
|
||||
|
||||
# Profile
|
||||
if provider == "jit":
|
||||
torch.cuda.nvtx.range_push("jit")
|
||||
fn()
|
||||
torch.cuda.nvtx.range_pop()
|
||||
elif provider == "sgl_kernel":
|
||||
torch.cuda.nvtx.range_push("sgl_kernel")
|
||||
fn()
|
||||
torch.cuda.nvtx.range_pop()
|
||||
|
||||
return run_benchmark(fn)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if not _SM100_SUPPORTED:
|
||||
print("[skip] MXFP8 MoE GroupMM benchmark requires sm100+ with CUDA 12.8+.")
|
||||
sys.exit(0)
|
||||
if not _SGL_KERNEL_AVAILABLE:
|
||||
print(f"[info] sgl-kernel baseline unavailable: {_SGL_KERNEL_REASON}")
|
||||
benchmark.run(print_data=True)
|
||||
@@ -1,77 +0,0 @@
|
||||
import itertools
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.testing
|
||||
|
||||
from sglang.kernels.jit.benchmark.utils import (
|
||||
DEFAULT_DEVICE,
|
||||
get_benchmark_range,
|
||||
run_benchmark,
|
||||
)
|
||||
from sglang.kernels.ops.speculative.resolve_future_token_ids import (
|
||||
resolve_future_token_ids_cuda,
|
||||
)
|
||||
from sglang.srt.utils import get_compiler_backend
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
|
||||
register_cuda_ci(
|
||||
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
|
||||
)
|
||||
register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd")
|
||||
|
||||
SIZE_LIST = get_benchmark_range(
|
||||
full_range=[2**n for n in range(4, 16)], # 16 … 32K elements
|
||||
ci_range=[256, 4096],
|
||||
)
|
||||
|
||||
configs = list(itertools.product(SIZE_LIST))
|
||||
|
||||
|
||||
def _torch_resolve(input_ids, future_map):
|
||||
input_ids[:] = torch.where(
|
||||
input_ids < 0,
|
||||
future_map[torch.clamp(-input_ids, min=0)],
|
||||
input_ids,
|
||||
)
|
||||
|
||||
|
||||
_compiled_resolve = torch.compile(
|
||||
_torch_resolve, dynamic=True, backend=get_compiler_backend()
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["size"],
|
||||
x_vals=configs,
|
||||
line_arg="provider",
|
||||
line_vals=["jit", "torch_compile", "torch"],
|
||||
line_names=["SGL JIT Kernel", "torch.compile", "PyTorch"],
|
||||
styles=[("blue", "-"), ("green", "-."), ("red", "--")],
|
||||
ylabel="us",
|
||||
plot_name="resolve-future-token-ids-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(size: int, provider: str):
|
||||
map_size = 8192
|
||||
future_map = torch.randint(
|
||||
0, 50000, (map_size,), dtype=torch.int64, device=DEFAULT_DEVICE
|
||||
)
|
||||
input_ids = torch.randint(
|
||||
-map_size + 1, 50000, (size,), dtype=torch.int64, device=DEFAULT_DEVICE
|
||||
)
|
||||
|
||||
if provider == "jit":
|
||||
fn = lambda: resolve_future_token_ids_cuda(input_ids.clone(), future_map)
|
||||
elif provider == "torch_compile":
|
||||
fn = lambda: _compiled_resolve(input_ids.clone(), future_map)
|
||||
else:
|
||||
fn = lambda: _torch_resolve(input_ids.clone(), future_map)
|
||||
|
||||
return run_benchmark(fn)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run(print_data=True)
|
||||
@@ -1,6 +1,7 @@
|
||||
"""GPU-free import / registry / selector tests for ``sglang.kernels`` (RFC #29630)."""
|
||||
|
||||
import importlib
|
||||
import importlib.util
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
@@ -84,6 +85,20 @@ def test_specs_well_formed():
|
||||
assert sep == ":" and mod and attr, spec.target
|
||||
|
||||
|
||||
def test_internal_registry_target_modules_exist():
|
||||
for spec in K.registry.all_specs():
|
||||
module, _, _ = spec.target.partition(":")
|
||||
if module.startswith("sglang.kernels."):
|
||||
assert importlib.util.find_spec(module) is not None, spec.target
|
||||
|
||||
|
||||
def test_sparse_linear_attention_registry_targets_forward_kernel():
|
||||
spec = K.registry.get_backend(
|
||||
"diffusion.sparse_linear_attn_fwd", KernelBackend.TRITON
|
||||
)
|
||||
assert spec.target.endswith(":_attn_fwd")
|
||||
|
||||
|
||||
def test_single_backend_resolves_without_backend():
|
||||
assert K.select_kernel("gemm.fp8_scaled_mm").backend is KernelBackend.AOT
|
||||
|
||||
|
||||
@@ -1,136 +0,0 @@
|
||||
"""Test for the kpool top-k transform JIT kernel.
|
||||
|
||||
Ported from the former AOT sgl-kernel test (sgl-kernel/tests/test_topk.py).
|
||||
The kernel selects pool groups at pool granularity, expands each selected group
|
||||
to ``pool_size`` token indices, and optionally transforms those token indices
|
||||
through a page table or a ragged offset.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.moe.kpool_topk_transform import fast_kpool_topk_transform_fused
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
|
||||
def _ref_torch_kpool_transform_impl(
|
||||
score: torch.Tensor,
|
||||
lengths: torch.Tensor,
|
||||
pool_size: int,
|
||||
topk: int,
|
||||
page_table: Optional[torch.Tensor] = None,
|
||||
topk_indices_offset: Optional[torch.Tensor] = None,
|
||||
seq_lens: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
rows = score.shape[0]
|
||||
group_topk = topk // pool_size
|
||||
offsets = torch.arange(pool_size, dtype=torch.int32, device=score.device)
|
||||
out_cols = topk + (pool_size - 1 if seq_lens is not None else 0)
|
||||
out = torch.full((rows, out_cols), -1, dtype=torch.int32, device=score.device)
|
||||
for i in range(rows):
|
||||
length = int(lengths[i].item())
|
||||
valid_count = min(length, group_topk)
|
||||
write_pos = 0
|
||||
if valid_count == 0:
|
||||
token_ids = torch.empty((0,), dtype=torch.int32, device=score.device)
|
||||
elif length <= group_topk:
|
||||
selected = torch.arange(length, dtype=torch.int32, device=score.device)
|
||||
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
|
||||
else:
|
||||
selected = torch.topk(
|
||||
score[i, :length], group_topk, dim=-1, sorted=False
|
||||
).indices.to(torch.int32)
|
||||
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
|
||||
if token_ids.numel() > 0:
|
||||
if page_table is not None:
|
||||
token_ids = page_table[i, token_ids.long()].to(torch.int32)
|
||||
elif topk_indices_offset is not None:
|
||||
token_ids = token_ids + topk_indices_offset[i].to(torch.int32)
|
||||
write_pos = valid_count * pool_size
|
||||
out[i, :write_pos] = token_ids[:write_pos]
|
||||
if seq_lens is not None:
|
||||
tail_count = int(seq_lens[i].item()) % pool_size
|
||||
if tail_count > 0:
|
||||
raw_tail = length * pool_size + torch.arange(
|
||||
tail_count, dtype=torch.int32, device=score.device
|
||||
)
|
||||
if page_table is not None:
|
||||
tail = page_table[i, raw_tail.long()].to(torch.int32)
|
||||
elif topk_indices_offset is not None:
|
||||
tail = raw_tail + topk_indices_offset[i].to(torch.int32)
|
||||
else:
|
||||
tail = raw_tail
|
||||
out[i, write_pos : write_pos + tail_count] = tail
|
||||
return out
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"pool_size,group_topk",
|
||||
[(16, 128), (16, 160), (16, 192), (16, 224), (8, 256), (4, 512)],
|
||||
)
|
||||
@pytest.mark.parametrize("mode", ["raw", "paged", "ragged"])
|
||||
@pytest.mark.parametrize("append_tail", [False, True])
|
||||
@torch.inference_mode()
|
||||
def test_kpool_topk_transform_kernel(
|
||||
pool_size: int, group_topk: int, mode: str, append_tail: bool
|
||||
) -> None:
|
||||
torch.manual_seed(42)
|
||||
bs = 17
|
||||
topk = pool_size * group_topk
|
||||
num_groups = 4096
|
||||
score = torch.randn(bs, num_groups, dtype=torch.float32, device="cuda")
|
||||
lengths = torch.randint(
|
||||
group_topk + 1, num_groups + 1, (bs,), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
page_table = None
|
||||
topk_indices_offset = None
|
||||
seq_lens = None
|
||||
tail_counts = torch.randint(0, pool_size, (bs,), dtype=torch.int32, device="cuda")
|
||||
if append_tail:
|
||||
seq_lens = lengths * pool_size + tail_counts
|
||||
if mode == "paged":
|
||||
page_table = torch.arange(
|
||||
bs * (num_groups * pool_size + pool_size),
|
||||
dtype=torch.int32,
|
||||
device="cuda",
|
||||
).view(bs, num_groups * pool_size + pool_size)
|
||||
elif mode == "ragged":
|
||||
topk_indices_offset = torch.randint(
|
||||
0, 2048, (bs,), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
out_ref = _ref_torch_kpool_transform_impl(
|
||||
score,
|
||||
lengths,
|
||||
pool_size,
|
||||
topk,
|
||||
page_table=page_table,
|
||||
topk_indices_offset=topk_indices_offset,
|
||||
seq_lens=seq_lens,
|
||||
)
|
||||
out_our = fast_kpool_topk_transform_fused(
|
||||
score,
|
||||
lengths,
|
||||
pool_size,
|
||||
topk,
|
||||
page_table=page_table,
|
||||
topk_indices_offset=topk_indices_offset,
|
||||
seq_lens=seq_lens,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
out_ref = torch.sort(out_ref, dim=-1).values
|
||||
out_our = torch.sort(out_our, dim=-1).values
|
||||
torch.testing.assert_close(out_our, out_ref, atol=0, rtol=0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
@@ -1,152 +0,0 @@
|
||||
import random
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.quantization.mxfp8 import (
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant,
|
||||
es_sm100_mxfp8_blockscaled_moe_grouped_gemm,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
|
||||
def align(val: int, alignment: int = 128) -> int:
|
||||
return int((val + alignment - 1) // alignment * alignment)
|
||||
|
||||
|
||||
# Copy from: https://github.com/deepseek-ai/DeepGEMM/blob/main/deep_gemm/utils.py
|
||||
def calc_diff(x, y):
|
||||
x, y = x.double(), y.double()
|
||||
denominator = (x * x + y * y).sum()
|
||||
sim = 2 * (x * y).sum() / denominator
|
||||
return 1 - sim
|
||||
|
||||
|
||||
def is_sm100_supported(device=None) -> bool:
|
||||
return (torch.cuda.get_device_capability(device)[0] == 10) and (
|
||||
torch.version.cuda >= "12.8"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_sm100_supported(),
|
||||
reason="test_mxfp8_moe at jit kernen is only supported on sm100",
|
||||
)
|
||||
@pytest.mark.parametrize("num_experts", [8, 16, 32, 64])
|
||||
@pytest.mark.parametrize("out_dtype", [torch.half, torch.bfloat16])
|
||||
def test_es_sm100_mxfp8_blockscaled_grouped_mm(num_experts, out_dtype):
|
||||
device = "cuda"
|
||||
alignment = 128
|
||||
n_g = random.randint(1, 64) * alignment
|
||||
k_g = random.randint(1, 64) * alignment
|
||||
|
||||
expert_offset = 0
|
||||
expert_offsets = []
|
||||
aux_expert_offset = 0
|
||||
aux_expert_offsets = []
|
||||
a_blockscale_offset = 0
|
||||
a_blockscale_offsets = []
|
||||
b_blockscale_offset = 0
|
||||
b_blockscale_offsets = []
|
||||
a_list = []
|
||||
b_list = []
|
||||
ref_d_list = []
|
||||
tokens_per_expert = []
|
||||
|
||||
for g in range(num_experts):
|
||||
m_g = random.randint(1, 512)
|
||||
tokens_per_expert.append(m_g)
|
||||
expert_offsets.append(expert_offset)
|
||||
expert_offset += m_g
|
||||
aux_expert_offsets.append(aux_expert_offset)
|
||||
aux_expert_offset += n_g
|
||||
a_blockscale_offsets.append(a_blockscale_offset)
|
||||
a_blockscale_offset += align(m_g, 128)
|
||||
b_blockscale_offsets.append(b_blockscale_offset)
|
||||
b_blockscale_offset += n_g # n_g already align to 128
|
||||
|
||||
a = torch.normal(
|
||||
0.0, std=1.0, size=(m_g, k_g), device=device, dtype=out_dtype
|
||||
) # (M, K):(K, 1)
|
||||
b = torch.normal(
|
||||
0.0, std=1.0, size=(n_g, k_g), device=device, dtype=out_dtype
|
||||
) # (N, K):(K, 1)
|
||||
|
||||
a_list.append(a)
|
||||
b_list.append(b)
|
||||
ref_d = a @ b.T
|
||||
ref_d_list.append(ref_d)
|
||||
a = torch.concat(a_list, dim=0)
|
||||
b = torch.concat(b_list, dim=0)
|
||||
|
||||
_expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32)
|
||||
_aux_expert_offsets = torch.tensor(aux_expert_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
_b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
|
||||
a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device)
|
||||
a_scale_factor = torch.zeros(
|
||||
(a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device
|
||||
)
|
||||
|
||||
b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device)
|
||||
b_scale_factor = torch.zeros(
|
||||
(num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device
|
||||
)
|
||||
tokens_per_expert = torch.tensor(tokens_per_expert).to(
|
||||
device=device, dtype=torch.int32
|
||||
)
|
||||
workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device)
|
||||
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant(
|
||||
a,
|
||||
tokens_per_expert,
|
||||
_expert_offsets,
|
||||
_a_blockscale_offsets,
|
||||
a_quant,
|
||||
a_scale_factor,
|
||||
)
|
||||
es_sm100_mxfp8_blockscaled_grouped_quant(
|
||||
b,
|
||||
torch.ones_like(tokens_per_expert) * n_g,
|
||||
_aux_expert_offsets,
|
||||
_b_blockscale_offsets,
|
||||
b_quant,
|
||||
b_scale_factor,
|
||||
)
|
||||
|
||||
b_quant = b_quant.view(num_experts, n_g, k_g)
|
||||
b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32)
|
||||
d = es_sm100_mxfp8_blockscaled_moe_grouped_gemm(
|
||||
b_quant,
|
||||
a_quant,
|
||||
b_scale_factor,
|
||||
a_scale_factor,
|
||||
_expert_offsets,
|
||||
_a_blockscale_offsets,
|
||||
tokens_per_expert,
|
||||
workspace,
|
||||
a.dtype,
|
||||
)
|
||||
|
||||
for g in range(num_experts):
|
||||
baseline = ref_d_list[g]
|
||||
actual = d[expert_offsets[g] : (expert_offsets[g] + tokens_per_expert[g])]
|
||||
diff = calc_diff(actual, baseline)
|
||||
assert diff < 0.001
|
||||
print(
|
||||
f"m_g={baseline.shape[0]} n_g={n_g} k_g={k_g} num_experts={num_experts}, out_dtype={out_dtype}, diff={diff:.5f}: OK"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
@@ -1,71 +0,0 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.speculative.resolve_future_token_ids import (
|
||||
resolve_future_token_ids_cuda,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=9, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
register_amd_ci(est_time=9, stage="jit-kernel-unit", runner_config="amd")
|
||||
|
||||
|
||||
def _reference_resolve(input_ids, future_map):
|
||||
"""Reference implementation using plain torch."""
|
||||
result = input_ids.clone()
|
||||
result[:] = torch.where(
|
||||
result < 0,
|
||||
future_map[torch.clamp(-result, min=0)],
|
||||
result,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@pytest.mark.parametrize("size", [1, 2, 127, 128, 255, 256, 1024, 4097])
|
||||
@pytest.mark.parametrize("dtype", [torch.int32, torch.int64])
|
||||
class TestResolveFutureTokenIds:
|
||||
def test_all_negative(self, size: int, dtype: torch.dtype) -> None:
|
||||
map_size = 8192
|
||||
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
|
||||
# Negative indices in range [-map_size+1, -1]
|
||||
input_ids = -torch.randint(1, map_size, (size,), dtype=dtype, device="cuda")
|
||||
|
||||
expected = _reference_resolve(input_ids, future_map)
|
||||
resolve_future_token_ids_cuda(input_ids, future_map)
|
||||
assert torch.equal(input_ids, expected)
|
||||
|
||||
def test_all_non_negative(self, size: int, dtype: torch.dtype) -> None:
|
||||
map_size = 16
|
||||
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
|
||||
input_ids = torch.randint(0, 50000, (size,), dtype=dtype, device="cuda")
|
||||
|
||||
expected = input_ids.clone()
|
||||
resolve_future_token_ids_cuda(input_ids, future_map)
|
||||
assert torch.equal(input_ids, expected)
|
||||
|
||||
def test_mixed(self, size: int, dtype: torch.dtype) -> None:
|
||||
map_size = 8192
|
||||
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
|
||||
# Mix of negative and non-negative
|
||||
input_ids = torch.randint(
|
||||
-map_size + 1, 50000, (size,), dtype=dtype, device="cuda"
|
||||
)
|
||||
|
||||
expected = _reference_resolve(input_ids, future_map)
|
||||
resolve_future_token_ids_cuda(input_ids, future_map)
|
||||
assert torch.equal(input_ids, expected)
|
||||
|
||||
def test_zeros(self, size: int, dtype: torch.dtype) -> None:
|
||||
map_size = 16
|
||||
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
|
||||
input_ids = torch.zeros(size, dtype=dtype, device="cuda")
|
||||
|
||||
expected = input_ids.clone()
|
||||
resolve_future_token_ids_cuda(input_ids, future_map)
|
||||
assert torch.equal(input_ids, expected)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
Reference in New Issue
Block a user