From 7d953440ec965032b6b5af386d56c427ae4e6da1 Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Fri, 20 Feb 2026 21:01:05 +0800 Subject: [PATCH] [jit kernel] Support per_token_group_quant_8bit jit kernel (#18905) Co-authored-by: luoyuan.luo --- .../bench_per_token_group_quant_8bit.py | 290 ++++++++++++++++++ .../csrc/gemm/per_token_group_quant_8bit.cuh | 218 +++++++++++++ .../jit_kernel/per_token_group_quant_8bit.py | 67 ++++ .../tests/test_per_token_group_quant_8bit.py | 205 +++++++++++++ python/sglang/jit_kernel/utils.py | 3 +- .../srt/layers/quantization/fp8_kernel.py | 42 ++- 6 files changed, 811 insertions(+), 14 deletions(-) create mode 100644 python/sglang/jit_kernel/benchmark/bench_per_token_group_quant_8bit.py create mode 100644 python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh create mode 100644 python/sglang/jit_kernel/per_token_group_quant_8bit.py create mode 100644 python/sglang/jit_kernel/tests/test_per_token_group_quant_8bit.py diff --git a/python/sglang/jit_kernel/benchmark/bench_per_token_group_quant_8bit.py b/python/sglang/jit_kernel/benchmark/bench_per_token_group_quant_8bit.py new file mode 100644 index 000000000..7edcf82c2 --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/bench_per_token_group_quant_8bit.py @@ -0,0 +1,290 @@ +import itertools +import os +from typing import Any, Dict, List + +import torch +import triton +from sgl_kernel.test_utils import create_per_token_group_quant_test_data + +from sglang.jit_kernel.benchmark.utils import ( + get_benchmark_range, +) +from sglang.jit_kernel.per_token_group_quant_8bit import ( + per_token_group_quant_8bit as sglang_per_token_group_quant_8bit, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + create_per_token_group_quant_fp8_output_scale, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + per_token_group_quant_8bit as triton_per_token_group_quant_8bit, +) +from sglang.srt.utils import is_hip +from sglang.srt.utils.bench_utils import bench_kineto + +# CI environment detection +IS_CI = ( + os.getenv("CI", "false").lower() == "true" + or os.getenv("GITHUB_ACTIONS", "false").lower() == "true" +) + +_is_hip = is_hip() +fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn + +NUM_TESTS = 300 if IS_CI else 30 + +GROUP_SIZE_RANGE = [128] +DST_DTYPE_RANGE = [fp8_type_] + +# ---- GEMM-like branch (num_ranks=None) ---- +NUM_TOKENS_RANGE_GEMM = get_benchmark_range( + full_range=[1, 4, 16, 64, 256, 768, 2048, 8192, 16384], + ci_range=[768], +) +HIDDEN_DIM_RANGE_GEMM = [1536, 7168, 16384] +NUM_RANKS_RANGE_GEMM = [None] + + +FLAGS_GEMM_FULL: List[Dict[str, Any]] = [ + dict( + column_major_scales=False, + scale_tma_aligned=False, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=False, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), +] +FLAGS_GEMM_CI: List[Dict[str, Any]] = [ + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), +] +FLAGS_RANGE_GEMM = get_benchmark_range( + full_range=FLAGS_GEMM_FULL, ci_range=FLAGS_GEMM_CI +) + +CONFIGS_GEMM = list( + itertools.product( + NUM_TOKENS_RANGE_GEMM, + HIDDEN_DIM_RANGE_GEMM, + GROUP_SIZE_RANGE, + NUM_RANKS_RANGE_GEMM, + DST_DTYPE_RANGE, + FLAGS_RANGE_GEMM, + ) +) + +# ---- MoE-like / multi-rank branch (hidden_dim=2048, num_ranks in {8,16,32,48}) ---- +NUM_TOKENS_RANGE_MOE = get_benchmark_range( + full_range=[1 * 8, 4 * 8, 64 * 8, 256 * 8, 768 * 8], + ci_range=[768 * 8], +) +HIDDEN_DIM_RANGE_MOE = [2048] +NUM_RANKS_RANGE_MOE = get_benchmark_range( + full_range=[8, 16, 32, 48], + ci_range=[48], +) + +FLAGS_MOE: List[Dict[str, Any]] = [ + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="balanced", + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="imbalanced", + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="extreme", + ), +] +FLAGS_RANGE_MOE = get_benchmark_range(full_range=FLAGS_MOE, ci_range=FLAGS_MOE) + +CONFIGS_MOE = list( + itertools.product( + NUM_TOKENS_RANGE_MOE, + HIDDEN_DIM_RANGE_MOE, + GROUP_SIZE_RANGE, + NUM_RANKS_RANGE_MOE, + DST_DTYPE_RANGE, + FLAGS_RANGE_MOE, + ) +) + +# ---- Final configs ---- +CONFIGS = CONFIGS_GEMM + CONFIGS_MOE + +LINE_VALS = ["triton", "sglang"] +LINE_NAMES = ["Triton (Inaccurate)", "SGL Kernel"] +STYLES = [("blue", "-"), ("green", "-")] + + +def _flatten_to_2d(t: torch.Tensor) -> torch.Tensor: + """Reshape a tensor with 3+ dims to 2D by merging all leading dims.""" + if t.ndim <= 2: + return t + return t.reshape(-1, t.shape[-1]) + + +def _make_sglang_bench_fn( + x: torch.Tensor, + group_size: int, + dst_dtype: torch.dtype, + flags: dict, +): + """ + Adapter that pre-allocates output tensors and returns a zero-arg callable + matching the JIT kernel's signature. + + The JIT kernel does not support fuse_silu_and_mul, so when enabled we + pre-compute silu+mul on the input. bench_kineto only times the kernel + matching the given name, so the pre-processing is not included. + + The JIT kernel expects 2D tensors, so any higher-dimensional inputs + (e.g. from masked_layout_mode) are flattened to 2D. + """ + fuse_silu_and_mul = flags.get("fuse_silu_and_mul", False) + column_major_scales = flags.get("column_major_scales", False) + scale_tma_aligned = flags.get("scale_tma_aligned", False) + scale_ue8m0 = flags.get("scale_ue8m0", False) + + # JIT kernel does not support fuse_silu_and_mul; pre-compute it + if fuse_silu_and_mul: + half = x.shape[-1] // 2 + x_input = torch.nn.functional.silu(x[..., :half]) * x[..., half:] + else: + x_input = x + + # JIT kernel expects 2D (num_tokens, hidden_dim); flatten if needed + x_input = _flatten_to_2d(x_input.contiguous()) + + out_shape = x_input.shape + output_q = torch.empty(out_shape, device=x.device, dtype=dst_dtype) + + fp8_max = torch.finfo(dst_dtype).max + fp8_min = -fp8_max + + output_s = create_per_token_group_quant_fp8_output_scale( + x_shape=out_shape, + device=x.device, + group_size=group_size, + column_major_scales=column_major_scales, + scale_tma_aligned=scale_tma_aligned, + scale_ue8m0=scale_ue8m0, + ) + + def _run(): + sglang_per_token_group_quant_8bit( + input=x_input, + output_q=output_q, + output_s=output_s, + group_size=group_size, + eps=1e-10, + fp8_min=fp8_min, + fp8_max=fp8_max, + scale_ue8m0=scale_ue8m0, + ) + + return _run + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=[ + "num_tokens", + "hidden_dim", + "group_size", + "num_ranks", + "dst_dtype", + "flags", + ], + x_vals=CONFIGS, + line_arg="provider", + line_vals=LINE_VALS, + # Triton has multi kernels and we only report the time for the core one + line_names=LINE_NAMES, + styles=STYLES, + ylabel="us", + plot_name="per-token-group-quant-8bit-performance", + args={}, + ) +) +def benchmark( + num_tokens, hidden_dim, group_size, num_ranks, dst_dtype, flags, provider +): + print( + f"Testing: {num_tokens=} {hidden_dim=} {group_size=} {num_ranks=} {dst_dtype=} {flags=} {provider=}" + ) + + x, masked_m = create_per_token_group_quant_test_data( + num_tokens=num_tokens, hidden_dim=hidden_dim, num_ranks=num_ranks, flags=flags + ) + + if provider == "triton": + fn = triton_per_token_group_quant_8bit + kernel_names = "_per_token_group_quant_8bit|_silu_and_mul_post_quant_kernel" + bench_fn = lambda: fn( + x=x, + masked_m=masked_m, + group_size=group_size, + dst_dtype=dst_dtype, + **{k: v for k, v in flags.items() if k not in ["masked_layout_mode"]}, + ) + elif provider == "sglang": + kernel_names = "per_token_group_quant_8bit_kernel" + bench_fn = _make_sglang_bench_fn( + x=x, + group_size=group_size, + dst_dtype=dst_dtype, + flags=flags, + ) + else: + raise ValueError(f"Unknown provider: {provider}") + + time_s = bench_kineto(bench_fn, kernel_names=kernel_names, num_tests=NUM_TESTS) + return time_s * 1e6 + + +if __name__ == "__main__": + benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh b/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh new file mode 100644 index 000000000..20724c92b --- /dev/null +++ b/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh @@ -0,0 +1,218 @@ +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +namespace { + +constexpr int kThreadsPerGroup = 16; + +__device__ __forceinline__ float GroupReduceMax(float val, const int tid) { + unsigned mask = threadIdx.x % 32 >= 16 ? 0xffff0000 : 0x0000ffff; + val = fmaxf(val, __shfl_xor_sync(mask, val, 8)); + val = fmaxf(val, __shfl_xor_sync(mask, val, 4)); + val = fmaxf(val, __shfl_xor_sync(mask, val, 2)); + val = fmaxf(val, __shfl_xor_sync(mask, val, 1)); + return val; +} + +template +using scale_packed_t_t = std::conditional_t; + +template +using scale_element_t_t = std::conditional_t; + +template +__global__ void per_token_group_quant_8bit_kernel( + const T* __restrict__ input, + DST_DTYPE* __restrict__ output_q, + scale_packed_t_t* __restrict__ output_s, + const int group_size, + const int num_groups, + const int groups_per_block, + const float eps, + const float min_8bit, + const float max_8bit, + const int num_groups_per_row = 0, + const int scale_stride = 0) { + using namespace device; + namespace math = device::math; + + (void)num_groups; + + const int local_group_id = static_cast(threadIdx.x / kThreadsPerGroup); + const int lane_id = threadIdx.x % kThreadsPerGroup; + + const int64_t block_group_id = blockIdx.x * groups_per_block; + const int64_t global_group_id = block_group_id + local_group_id; + const int64_t block_group_offset = global_group_id * group_size; + + float local_absmax = eps; + + using scale_packed_t = scale_packed_t_t; + using scale_element_t = scale_element_t_t; + static_assert(sizeof(scale_packed_t) % sizeof(scale_element_t) == 0); + + const T* group_input = input + block_group_offset; + DST_DTYPE* group_output = static_cast(output_q) + block_group_offset; + scale_element_t* scale_output = nullptr; + + if constexpr (kIsColumnMajor) { + constexpr int kElemsPerPack = static_cast(sizeof(scale_packed_t) / sizeof(scale_element_t)); + const int row_idx = global_group_id / num_groups_per_row; + const int col_idx_unpacked = global_group_id % num_groups_per_row; + const int col_idx = col_idx_unpacked / kElemsPerPack; + const int pack_idx = col_idx_unpacked % kElemsPerPack; + scale_output = reinterpret_cast(output_s) + + (col_idx * scale_stride * kElemsPerPack + row_idx * kElemsPerPack + pack_idx); + } else { + static_assert(!kScaleUE8M0); + scale_output = output_s + global_group_id; + } + + constexpr uint32_t kVecSize = 16 / sizeof(T); + using vec_t = AlignedVector; + const auto gmem_in = tile::Memory::thread(); + + const int32_t num_vec_elems = group_size / kVecSize; + + for (int32_t i = lane_id; i < num_vec_elems; i += kThreadsPerGroup) { + const vec_t input_vec = gmem_in.load(group_input, i); + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const float val = static_cast(input_vec[j]); + local_absmax = math::max(local_absmax, math::abs(val)); + } + } + + local_absmax = GroupReduceMax(local_absmax, lane_id); + + float y_s = local_absmax / max_8bit; + if constexpr (kScaleUE8M0) { + y_s = exp2f(ceilf(log2f(math::max(y_s, 1e-10f)))); + } + + scale_element_t y_s_quant; + if constexpr (kScaleUE8M0) { + y_s_quant = static_cast(((int)log2f(y_s)) + 127); + } else { + y_s_quant = y_s; + } + + if (lane_id == 0) { + *scale_output = y_s_quant; + } + + for (int32_t i = lane_id; i < num_vec_elems; i += kThreadsPerGroup) { + const vec_t input_vec = gmem_in.load(group_input, i); + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const float val = static_cast(input_vec[j]); + const float q_val = math::min(math::max(val / y_s, min_8bit), max_8bit); + group_output[i * kVecSize + j] = DST_DTYPE(q_val); + } + } +} + +inline int compute_groups_per_block(int64_t num_groups) { + if (num_groups % 16 == 0) return 16; + if (num_groups % 8 == 0) return 8; + if (num_groups % 4 == 0) return 4; + if (num_groups % 2 == 0) return 2; + return 1; +} + +template +void per_token_group_quant_8bit( + tvm::ffi::TensorView input, + tvm::ffi::TensorView output_q, + tvm::ffi::TensorView output_s, + int64_t group_size, + double eps, + double min_8bit, + double max_8bit, + bool scale_ue8m0) { + using namespace host; + + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto K = SymbolicSize{"hidden_dim"}; + device.set_options(); + + TensorMatcher({M, K}).with_dtype().with_device(device).verify(input); + TensorMatcher({M, K}).with_dtype().with_device(device).verify(output_q); + + const auto num_tokens = M.unwrap(); + const auto hidden_dim = K.unwrap(); + + const int64_t num_groups_per_row = hidden_dim / group_size; + const int64_t num_groups = num_tokens * num_groups_per_row; + + const int groups_per_block = compute_groups_per_block(num_groups); + const int num_blocks = num_groups / groups_per_block; + const int num_threads = groups_per_block * kThreadsPerGroup; + const bool is_column_major = output_s.stride(0) < output_s.stride(1); + const int scale_stride = output_s.stride(1); + + const float feps = static_cast(eps); + const float fmin8 = static_cast(min_8bit); + const float fmax8 = static_cast(max_8bit); + + if (is_column_major) { + if (scale_ue8m0) { + LaunchKernel(num_blocks, num_threads, input.device())( + per_token_group_quant_8bit_kernel, + static_cast(input.data_ptr()), + static_cast(output_q.data_ptr()), + static_cast(output_s.data_ptr()), + static_cast(group_size), + static_cast(num_groups), + static_cast(groups_per_block), + feps, + fmin8, + fmax8, + static_cast(num_groups_per_row), + scale_stride); + } else { + LaunchKernel(num_blocks, num_threads, input.device())( + per_token_group_quant_8bit_kernel, + static_cast(input.data_ptr()), + static_cast(output_q.data_ptr()), + static_cast(output_s.data_ptr()), + static_cast(group_size), + static_cast(num_groups), + static_cast(groups_per_block), + feps, + fmin8, + fmax8, + static_cast(num_groups_per_row), + scale_stride); + } + } else { + LaunchKernel(num_blocks, num_threads, input.device())( + per_token_group_quant_8bit_kernel, + static_cast(input.data_ptr()), + static_cast(output_q.data_ptr()), + static_cast(output_s.data_ptr()), + static_cast(group_size), + static_cast(num_groups), + static_cast(groups_per_block), + feps, + fmin8, + fmax8, + 0, + 0); + } +} +} // namespace diff --git a/python/sglang/jit_kernel/per_token_group_quant_8bit.py b/python/sglang/jit_kernel/per_token_group_quant_8bit.py new file mode 100644 index 000000000..701e37820 --- /dev/null +++ b/python/sglang/jit_kernel/per_token_group_quant_8bit.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +from sglang.jit_kernel.utils import CPP_DTYPE_MAP as OUTPUT_DTYPE_MAP + + +@cache_once +def _jit_per_token_group_quant_8bit_module( + dtype: torch.dtype, output_type: torch.dtype +) -> Module: + input_args = make_cpp_args(dtype) + out_cpp = OUTPUT_DTYPE_MAP[output_type] + return load_jit( + "per_token_group_quant_8bit", + cuda_files=["gemm/per_token_group_quant_8bit.cuh"], + cuda_wrappers=[ + ( + "per_token_group_quant_8bit", + f"per_token_group_quant_8bit<{input_args}, {out_cpp}>", + ) + ], + ) + + +def per_token_group_quant_8bit( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + eps: float, + fp8_min: float, + fp8_max: float, + scale_ue8m0: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + """ + Per-token-group quantization to 8-bit format. + + Args: + input: Input tensor to quantize (float, half, or bfloat16). + output_q: Output quantized tensor (e.g., fp8_e4m3 or int8). + output_s: Output scale tensor. + group_size: The size of the group for quantization. + eps: A small value to avoid division by zero. + fp8_min: The minimum value of the 8-bit data type. + fp8_max: The maximum value of the 8-bit data type. + scale_ue8m0: Whether to use UE8M0 format for scales. + """ + module = _jit_per_token_group_quant_8bit_module(input.dtype, output_q.dtype) + module.per_token_group_quant_8bit( + input, + output_q, + output_s, + group_size, + eps, + fp8_min, + fp8_max, + scale_ue8m0, + ) + return output_q, output_s diff --git a/python/sglang/jit_kernel/tests/test_per_token_group_quant_8bit.py b/python/sglang/jit_kernel/tests/test_per_token_group_quant_8bit.py new file mode 100644 index 000000000..55f6f35be --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_per_token_group_quant_8bit.py @@ -0,0 +1,205 @@ +import itertools + +import pytest +import torch + +from sglang.srt.utils import is_hip + +_is_hip = is_hip() +fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn + +from sgl_kernel.test_utils import ( + assert_all_close_or_tiny_diff, + create_per_token_group_quant_test_data, +) + +from sglang.jit_kernel.per_token_group_quant_8bit import ( + per_token_group_quant_8bit as sglang_per_token_group_quant_8bit, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + create_per_token_group_quant_fp8_output_scale, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + per_token_group_quant_8bit as triton_per_token_group_quant_8bit, +) + +configs = list( + itertools.product( + [1, 4, 16, 64, 127, 128, 512, 1024, 4096, 8192], # num_tokens + [128, 256, 384, 512, 1024, 1536, 1664, 2048, 4096, 7168, 16384], # hidden_dim + [16, 32, 64, 128], # group_size + [None], # num_ranks + [fp8_type_], # dtype + [ + dict( + column_major_scales=False, + scale_tma_aligned=False, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=False, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=False, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=False, + masked_layout_mode=None, + ), + ], + ) +) + list( + itertools.product( + [1, 4, 1 * 8, 4 * 8, 64 * 8, 256 * 8, 768 * 8], + [2048], + [128], + [8, 16, 32, 48], + [fp8_type_], + [ + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode=None, + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="balanced", + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="imbalanced", + ), + dict( + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_layout_mode="extreme", + ), + ], + ) +) + + +@pytest.mark.parametrize( + "num_tokens, hidden_dim, group_size, num_ranks, dst_dtype, flags", configs +) +def test_per_token_group_quant_with_column_major( + num_tokens, + hidden_dim, + group_size, + num_ranks, + dst_dtype, + flags, +): + arch_major, _ = torch.cuda.get_device_capability(torch.cuda.current_device()) + if flags["scale_ue8m0"] and (arch_major <= 9): + pytest.skip("Only Blackwell need ue8m0 fusion") + return + + if (flags["scale_ue8m0"] and (group_size != 128)) or ( + (dst_dtype == torch.int8) and flags["column_major_scales"] + ): + pytest.skip() + return + + x, masked_m = create_per_token_group_quant_test_data( + num_tokens=num_tokens, hidden_dim=hidden_dim, num_ranks=num_ranks, flags=flags + ) + + execute_kwargs = dict( + x=x, + masked_m=masked_m, + group_size=group_size, + eps=1e-10, + dst_dtype=dst_dtype, + **{k: v for k, v in flags.items() if k not in ["masked_layout_mode"]}, + ) + + def _postprocess(x_q, x_s): + if masked_m is not None: + print(f"Mask tokens after {masked_m} to be zero") + for i in range(len(masked_m)): + x_q[i, masked_m[i] :, :] = 0 + x_s[i, masked_m[i] :, :] = 0 + return x_q, x_s + + x_q_triton, x_s_triton = _postprocess( + *triton_per_token_group_quant_8bit(**execute_kwargs) + ) + + fuse_silu_and_mul = False + out_shape = (*x.shape[:-1], x.shape[-1] // (2 if fuse_silu_and_mul else 1)) + + fp8_dtype = torch.float8_e4m3fn + fp8_max = torch.finfo(fp8_dtype).max + fp8_min = -fp8_max + x_q = torch.empty(out_shape, device=x.device, dtype=fp8_dtype) + x_s = create_per_token_group_quant_fp8_output_scale( + x_shape=out_shape, + device=x.device, + group_size=group_size, + column_major_scales=False, + scale_tma_aligned=False, + scale_ue8m0=False, + ) + + execute_kwargs = dict( + input=x, + output_q=x_q, + output_s=x_s, + group_size=group_size, + eps=1e-10, + fp8_max=fp8_max, + fp8_min=fp8_min, + ) + x_q_sglang, x_s_sglang = _postprocess( + *sglang_per_token_group_quant_8bit(**execute_kwargs) + ) + + try: + assert_all_close_or_tiny_diff(x_q_triton, x_q_sglang) + torch.testing.assert_close( + x_s_triton.contiguous(), + x_s_sglang.contiguous(), + rtol=1e-3, + atol=1e-5, + msg=lambda message: message + f" {x_s_triton=} {x_s_sglang=}", + ) + except AssertionError: + print( + f"{x.shape=} {x_q_triton.shape=} {x_s_triton.shape=} {x_q_sglang.shape=} {x_s_sglang.shape=}" + ) + print(f"{x=}") + print(f"{masked_m=}") + print(f"{x_q_triton=}") + print(f"{x_s_triton=}") + print(f"{x_q_sglang=}") + print(f"{x_s_sglang=}") + + raise + + +if __name__ == "__main__": + pytest.main([__file__]) diff --git a/python/sglang/jit_kernel/utils.py b/python/sglang/jit_kernel/utils.py index ede47831a..04d49fcdc 100644 --- a/python/sglang/jit_kernel/utils.py +++ b/python/sglang/jit_kernel/utils.py @@ -10,7 +10,6 @@ import torch if TYPE_CHECKING: from tvm_ffi import Module - F = TypeVar("F", bound=Callable[..., Any]) @@ -73,7 +72,9 @@ class CPPArgList(list[str]): CPP_DTYPE_MAP = { torch.float: "fp32_t", torch.float16: "fp16_t", + torch.float8_e4m3fn: "fp8_e4m3_t", torch.bfloat16: "bf16_t", + torch.int8: "int8_t", } diff --git a/python/sglang/srt/layers/quantization/fp8_kernel.py b/python/sglang/srt/layers/quantization/fp8_kernel.py index 1806c54cd..4a143d724 100644 --- a/python/sglang/srt/layers/quantization/fp8_kernel.py +++ b/python/sglang/srt/layers/quantization/fp8_kernel.py @@ -63,6 +63,10 @@ if _is_cuda: enable_sgl_per_token_group_quant_8bit = False + from sglang.jit_kernel.per_token_group_quant_8bit import ( + per_token_group_quant_8bit as sgl_per_token_group_quant_8bit_jit, + ) + if _is_hip: _has_vllm = False if _use_aiter: @@ -501,19 +505,31 @@ def sglang_per_token_group_quant_fp8( if x.shape[0] > 0: # Temporary if enable_sgl_per_token_group_quant_8bit: - sgl_per_token_group_quant_8bit( - x, - x_q, - x_s, - group_size, - eps, - fp8_min, - fp8_max, - scale_ue8m0, - fuse_silu_and_mul, - masked_m, - enable_v2=enable_v2, - ) + if enable_v2: + sgl_per_token_group_quant_8bit( + x, + x_q, + x_s, + group_size, + eps, + fp8_min, + fp8_max, + scale_ue8m0, + fuse_silu_and_mul, + masked_m, + enable_v2=True, + ) + else: + sgl_per_token_group_quant_8bit_jit( + input=x, + output_q=x_q, + output_s=x_s, + group_size=group_size, + eps=eps, + fp8_min=fp8_min, + fp8_max=fp8_max, + scale_ue8m0=scale_ue8m0, + ) else: assert not enable_v2 sgl_per_token_group_quant_fp8(