diff --git a/python/sglang/kernels/aot/benchmark/bench_per_token_quant_fp8.py b/python/sglang/kernels/aot/benchmark/bench_per_token_quant_fp8.py deleted file mode 100644 index dca014676..000000000 --- a/python/sglang/kernels/aot/benchmark/bench_per_token_quant_fp8.py +++ /dev/null @@ -1,228 +0,0 @@ -import itertools -import os -from typing import Optional, Tuple - -import torch -import triton -import triton.testing -from sgl_kernel import sgl_per_token_quant_fp8 - -from sglang.utils import is_in_ci - -# Optional vLLM import -try: - from vllm import _custom_ops as ops - - VLLM_AVAILABLE = True -except ImportError: - ops = None - VLLM_AVAILABLE = False - -from sglang.srt.utils import is_hip - -_is_hip = is_hip() - -IS_CI = is_in_ci() - -fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn - -# Get correct FP8 E4M3 maximum value -if _is_hip: - FP8_E4M3_MAX = 224.0 # ROCM uses 224.0 -else: - # For CUDA, get the actual max value from the type - FP8_E4M3_MAX = float(torch.finfo(fp8_type_).max) - - -def torch_per_token_quant_fp8( - input: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Pure PyTorch reference implementation for per-token FP8 quantization.""" - device = input.device - dtype = input.dtype - - # Find max absolute value per token (row) - exactly like CUDA kernel - max_vals = torch.abs(input).max(dim=1)[0] # [num_tokens] - - # Calculate scale per token - exactly like CUDA kernel: scale = max_value / FP8_E4M3_MAX - scales = max_vals / FP8_E4M3_MAX # [num_tokens] - - # No special zero handling - directly compute 1.0 / scale like CUDA kernel - scale_inv = 1.0 / scales # [num_tokens] - - # Quantize: input * scale_inv, then clamp to FP8 range - quantized_float = input * scale_inv.unsqueeze(1) # Broadcast scale_inv - quantized_float = torch.clamp(quantized_float, -FP8_E4M3_MAX, FP8_E4M3_MAX) - - # Convert to FP8 - use more explicit conversion - quantized_fp8 = quantized_float.to(fp8_type_) - - return quantized_fp8, scales - - -def vllm_per_token_quant_fp8( - input: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - if not VLLM_AVAILABLE: - # Fallback to SGLang implementation - return sglang_per_token_quant_fp8(input) - return ops.scaled_fp8_quant(input, use_per_token_if_dynamic=True) - - -def sglang_per_token_quant_fp8( - input: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - scale = torch.zeros(input.size(0), device=input.device, dtype=torch.float32) - output = torch.empty_like(input, device=input.device, dtype=fp8_type_) - sgl_per_token_quant_fp8(input, output, scale) - - return output, scale - - -def calculate_diff(batch_size: int, seq_len: int, hidden_dim: int): - """Compare Torch reference, VLLM, and SGLang implementations.""" - device = torch.device("cuda") - x = torch.rand( - (batch_size * seq_len, hidden_dim), dtype=torch.float16, device=device - ) - - # Get all three implementations - torch_out, torch_scale = torch_per_token_quant_fp8(x) - vllm_out, vllm_scale = vllm_per_token_quant_fp8(x) - sglang_out, sglang_scale = sglang_per_token_quant_fp8(x) - - if not VLLM_AVAILABLE: - print("⚠️ vLLM not available, skipping vLLM comparison") - # Only compare Torch vs SGLang - torch_sglang_scale_diff = torch.abs(torch_scale - sglang_scale).mean().item() - torch_sglang_out_diff = ( - torch.abs(torch_out.float() - sglang_out.float()).mean().item() - ) - print(f"Scale difference (Torch vs SGLang): {torch_sglang_scale_diff:.8f}") - print(f"Output difference (Torch vs SGLang): {torch_sglang_out_diff:.8f}") - return - - print(f"\n=== Comparison for hidden_dim={hidden_dim} ===") - - # Compare scales - torch_vllm_scale_diff = torch.abs(torch_scale - vllm_scale).mean().item() - torch_sglang_scale_diff = torch.abs(torch_scale - sglang_scale).mean().item() - vllm_sglang_scale_diff = torch.abs(vllm_scale - sglang_scale).mean().item() - - print(f"Scale differences:") - print(f" Torch vs VLLM: {torch_vllm_scale_diff:.8f}") - print(f" Torch vs SGLang: {torch_sglang_scale_diff:.8f}") - print(f" VLLM vs SGLang: {vllm_sglang_scale_diff:.8f}") - - # Compare outputs - torch_vllm_out_diff = torch.abs(torch_out.float() - vllm_out.float()).mean().item() - torch_sglang_out_diff = ( - torch.abs(torch_out.float() - sglang_out.float()).mean().item() - ) - vllm_sglang_out_diff = ( - torch.abs(vllm_out.float() - sglang_out.float()).mean().item() - ) - - print(f"Output differences:") - print(f" Torch vs VLLM: {torch_vllm_out_diff:.8f}") - print(f" Torch vs SGLang: {torch_sglang_out_diff:.8f}") - print(f" VLLM vs SGLang: {vllm_sglang_out_diff:.8f}") - - # Check tolerances - rtol, atol = 1e-3, 1e-5 - - torch_vllm_match = torch.allclose( - torch_out.float(), vllm_out.float(), rtol=rtol, atol=atol - ) and torch.allclose(torch_scale, vllm_scale, rtol=rtol, atol=atol) - torch_sglang_match = torch.allclose( - torch_out.float(), sglang_out.float(), rtol=rtol, atol=atol - ) and torch.allclose(torch_scale, sglang_scale, rtol=rtol, atol=atol) - - if hidden_dim == 1368: - rtol = 1e-2 - # we found vllm sglang has diff when hidden dim is not dividable by 16 - # and we believe SGLang is closer to Torch implementation - - vllm_sglang_match = torch.allclose( - vllm_out.float(), sglang_out.float(), rtol=rtol, atol=atol - ) and torch.allclose(vllm_scale, sglang_scale, rtol=rtol, atol=atol) - - print(f"Matches (rtol={rtol}, atol={atol}):") - print(f" Torch vs VLLM: {'✅' if torch_vllm_match else '❌'}") - print(f" Torch vs SGLang: {'✅' if torch_sglang_match else '❌'}") - print(f" VLLM vs SGLang: {'✅' if vllm_sglang_match else '❌'}") - - -# CI environment uses simplified parameters -if IS_CI: - batch_size_range = [16] # Single batch size for CI - seq_len_range = [64] # Single sequence length for CI - hidden_dim_range = [2048] # Single hidden dimension for CI -else: - batch_size_range = [16, 32, 64, 128] - seq_len_range = [64, 128, 256, 512, 1024, 2048, 4096] - hidden_dim_range = [1368, 2048, 4096] - -configs = list(itertools.product(batch_size_range, seq_len_range, hidden_dim_range)) - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["batch_size", "seq_len", "hidden_dim"], - x_vals=configs, - line_arg="provider", - line_vals=( - ["torch", "vllm", "sglang"] if VLLM_AVAILABLE else ["torch", "sglang"] - ), - line_names=( - ["Torch Reference", "VLLM", "SGL Kernel"] - if VLLM_AVAILABLE - else ["Torch Reference", "SGL Kernel"] - ), - styles=( - [("red", "-"), ("blue", "-"), ("green", "-")] - if VLLM_AVAILABLE - else [("red", "-"), ("green", "-")] - ), - ylabel="us", - plot_name="per-token-dynamic-quant-fp8-performance", - args={}, - ) -) -def benchmark_quantization(batch_size, seq_len, hidden_dim, provider): - dtype = torch.float16 - device = torch.device("cuda") - - x = torch.randn(batch_size * seq_len, hidden_dim, device=device, dtype=dtype) - - quantiles = [0.5, 0.2, 0.8] - - if provider == "torch": - fn = lambda: torch_per_token_quant_fp8(x.clone()) - elif provider == "vllm": - if not VLLM_AVAILABLE: - return (0, 0, 0) - fn = lambda: vllm_per_token_quant_fp8(x.clone()) - elif provider == "sglang": - fn = lambda: sglang_per_token_quant_fp8(x.clone()) - - ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(fn, quantiles=quantiles) - - return 1000 * ms, 1000 * max_ms, 1000 * min_ms - - -if __name__ == "__main__": - # Test various hidden dimensions for correctness - simplified for CI - if IS_CI: - test_dims = [2048] # Single dimension for CI - batch_size, seq_len = 4, 64 # Smaller values for CI - else: - test_dims = [1368, 2048, 4096] - batch_size, seq_len = 4, 4096 - - for dim in test_dims: - calculate_diff(batch_size=batch_size, seq_len=seq_len, hidden_dim=dim) - - print("\n" + "=" * 60) - print("Starting performance benchmark...") - benchmark_quantization.run(print_data=True) diff --git a/python/sglang/kernels/aot/csrc/common_extension.cc b/python/sglang/kernels/aot/csrc/common_extension.cc index 37c1a8fa8..4f6665cf0 100644 --- a/python/sglang/kernels/aot/csrc/common_extension.cc +++ b/python/sglang/kernels/aot/csrc/common_extension.cc @@ -133,6 +133,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { " float eps, float fp8_min, float fp8_max, bool scale_ue8m0, bool fuse_silu_and_mul, Tensor? masked_m) -> ()"); m.impl("sgl_per_token_group_quant_8bit_v2", torch::kCUDA, &sgl_per_token_group_quant_8bit_v2); + // Compatibility API: SGLang runtime dispatches to the JIT implementation, + // but external sgl_kernel consumers still rely on this exported CUDA op. m.def("sgl_per_token_quant_fp8(Tensor input, Tensor! output_q, Tensor! output_s) -> ()"); m.impl("sgl_per_token_quant_fp8", torch::kCUDA, &sgl_per_token_quant_fp8); diff --git a/python/sglang/kernels/aot/tests/test_per_token_quant_fp8.py b/python/sglang/kernels/aot/tests/test_per_token_quant_fp8.py deleted file mode 100644 index 9073936ca..000000000 --- a/python/sglang/kernels/aot/tests/test_per_token_quant_fp8.py +++ /dev/null @@ -1,64 +0,0 @@ -import itertools -import sys -from typing import Tuple - -import pytest -import torch -from sgl_kernel import sgl_per_token_quant_fp8 - -from sglang.srt.utils import is_hip - -_is_hip = is_hip() -fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn - - -def torch_per_token_quant_fp8(tensor, inv_scale): - # The reference implementation that fully aligns to - # the kernel being tested. - finfo = torch.finfo(torch.float8_e4m3fn) - inv_scale = inv_scale.view(-1, 1) - scale = inv_scale.reciprocal() - qweight = (tensor.to(torch.float32) * scale).clamp(min=finfo.min, max=finfo.max) - qweight = qweight.to(torch.float8_e4m3fn) - return qweight - - -def sglang_per_token_quant_fp8( - input: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - scale = torch.zeros(input.size(0), device=input.device, dtype=torch.float32) - output = torch.empty_like(input, device=input.device, dtype=fp8_type_) - - sgl_per_token_quant_fp8(input, output, scale) - scale = scale.reshape(-1, 1) - - return output, scale - - -PER_TOKEN_QUANT_CASES = list( - itertools.product([128, 256, 512], [512, 1076, 1368, 2048, 4096]) -) + [ - (39, 1536), - (1392, 1536), - (7807, 1536), -] - - -@pytest.mark.parametrize("num_tokens,hidden_dim", PER_TOKEN_QUANT_CASES) -def test_per_token_quant_compare_implementations( - num_tokens: int, - hidden_dim: int, -): - device = torch.device("cuda") - x = torch.rand((num_tokens, hidden_dim), dtype=torch.float16, device=device) - - sglang_out, sglang_scale = sglang_per_token_quant_fp8(x) - torch_out = torch_per_token_quant_fp8(x, sglang_scale) - - torch.testing.assert_close( - sglang_out.float(), torch_out.float(), rtol=1e-3, atol=1e-3 - ) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/python/sglang/kernels/jit/csrc/gemm/per_token_quant_fp8.cuh b/python/sglang/kernels/jit/csrc/gemm/per_token_quant_fp8.cuh new file mode 100644 index 000000000..140f7425a --- /dev/null +++ b/python/sglang/kernels/jit/csrc/gemm/per_token_quant_fp8.cuh @@ -0,0 +1,177 @@ +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +namespace sglang { + +constexpr uint32_t kPerTokenQuantWarpSize = 32; +constexpr uint32_t kPerTokenQuantTokensPerCTA = 8; + +/** \brief Quantize one token per warp for large token batches. */ +template +__global__ void per_token_quant_fp8_warp_kernel( + const T* __restrict__ input, + fp8_e4m3_t* __restrict__ output_q, + float* __restrict__ output_s, + uint32_t hidden_dim, + uint32_t num_tokens) { + using namespace device; + using input_vec_t = AlignedVector; + using output_vec_t = AlignedVector; + + const uint32_t warp_id = threadIdx.x / kPerTokenQuantWarpSize; + const uint32_t lane_id = threadIdx.x % kPerTokenQuantWarpSize; + const uint32_t token_id = blockIdx.x * kPerTokenQuantTokensPerCTA + warp_id; + if (token_id >= num_tokens) { + return; + } + + const T* token_input = input + token_id * hidden_dim; + fp8_e4m3_t* token_output = output_q + token_id * hidden_dim; + const uint32_t num_vecs = hidden_dim / kVecSize; + + float max_value = 0.0f; + for (uint32_t i = lane_id; i < num_vecs; i += kPerTokenQuantWarpSize) { + input_vec_t input_vec; + input_vec.load(token_input, i); +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + max_value = math::max(max_value, math::abs(static_cast(input_vec[j]))); + } + } + + const float scale = warp::reduce_max(max_value) / math::FP8_E4M3_MAX; + if (lane_id == 0) { + output_s[token_id] = scale; + } + const float scale_inv = scale == 0.0f ? 0.0f : 1.0f / scale; + + for (uint32_t i = lane_id; i < num_vecs; i += kPerTokenQuantWarpSize) { + input_vec_t input_vec; + output_vec_t output_vec; + input_vec.load(token_input, i); +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const float value = static_cast(input_vec[j]) * scale_inv; + output_vec[j] = static_cast(math::max(math::min(value, math::FP8_E4M3_MAX), -math::FP8_E4M3_MAX)); + } + output_vec.store(token_output, i); + } +} + +/** \brief Quantize one token per CTA for small token batches. */ +template +__global__ void per_token_quant_fp8_cta_kernel( + const T* __restrict__ input, fp8_e4m3_t* __restrict__ output_q, float* __restrict__ output_s, uint32_t hidden_dim) { + using namespace device; + using input_vec_t = AlignedVector; + using output_vec_t = AlignedVector; + + const uint32_t token_id = blockIdx.x; + const T* token_input = input + token_id * hidden_dim; + fp8_e4m3_t* token_output = output_q + token_id * hidden_dim; + const uint32_t num_vecs = hidden_dim / kVecSize; + + float max_value = 0.0f; + for (uint32_t i = threadIdx.x; i < num_vecs; i += blockDim.x) { + input_vec_t input_vec; + input_vec.load(token_input, i); +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + max_value = math::max(max_value, math::abs(static_cast(input_vec[j]))); + } + } + + __shared__ float reduction_smem[kPerTokenQuantWarpSize]; + __shared__ float scale_smem; + cta::reduce_max(max_value, reduction_smem); + __syncthreads(); + if (threadIdx.x == 0) { + scale_smem = reduction_smem[0] / math::FP8_E4M3_MAX; + output_s[token_id] = scale_smem; + } + __syncthreads(); + const float scale_inv = 1.0f / scale_smem; + + for (uint32_t i = threadIdx.x; i < num_vecs; i += blockDim.x) { + input_vec_t input_vec; + output_vec_t output_vec; + input_vec.load(token_input, i); +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const float value = static_cast(input_vec[j]) * scale_inv; + output_vec[j] = static_cast(math::max(math::min(value, math::FP8_E4M3_MAX), -math::FP8_E4M3_MAX)); + } + output_vec.store(token_output, i); + } +} + +template +void launch_per_token_quant_fp8( + DLDevice device, const T* input, fp8_e4m3_t* output_q, float* output_s, uint32_t hidden_dim, uint32_t num_tokens) { + constexpr uint32_t kBlockSize = 256; + const uint32_t sm_count = host::runtime::get_sm_count(device.device_id); + const bool use_warp_kernel = num_tokens >= sm_count * 2 * kPerTokenQuantTokensPerCTA; + if (use_warp_kernel) { + const uint32_t grid = host::div_ceil(num_tokens, kPerTokenQuantTokensPerCTA); + host::LaunchKernel(grid, kBlockSize, device)( + per_token_quant_fp8_warp_kernel, input, output_q, output_s, hidden_dim, num_tokens); + } else { + host::LaunchKernel(num_tokens, kBlockSize, device)( + per_token_quant_fp8_cta_kernel, input, output_q, output_s, hidden_dim); + } +} + +/** \brief Validate and launch dynamic per-token FP8 E4M3 quantization. */ +template +void per_token_quant_fp8(tvm::ffi::TensorView input, tvm::ffi::TensorView output_q, tvm::ffi::TensorView output_s) { + using namespace host; + auto M = SymbolicSize{"num_tokens"}; + auto MOutput = SymbolicSize{"output_num_tokens"}; + auto K = SymbolicSize{"hidden_dim"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({M, K}).with_dtype().with_device(device).verify(input); + TensorMatcher({MOutput, K}).with_dtype().with_device(device).verify(output_q); + TensorMatcher({MOutput, 1}).with_dtype().with_device(device).verify(output_s); + + CHECK_HOST(M.unwrap() > 0) << "per_token_quant_fp8: num_tokens must be positive"; + CHECK_HOST(MOutput.unwrap() >= M.unwrap()) + << "per_token_quant_fp8: output buffers must have at least " << M.unwrap() << " rows, got " << MOutput.unwrap(); + CHECK_HOST(K.unwrap() > 0 && K.unwrap() % 4 == 0) + << "per_token_quant_fp8: hidden_dim must be positive and divisible by 4, got " << K.unwrap(); + CHECK_HOST(M.unwrap() <= UINT32_MAX && K.unwrap() <= UINT32_MAX) + << "per_token_quant_fp8: dimensions exceed uint32 indexing"; + + const uint32_t num_tokens = static_cast(M.unwrap()); + const uint32_t hidden_dim = static_cast(K.unwrap()); + constexpr uint32_t kMaxVecSize = 16 / sizeof(T); + if (hidden_dim % kMaxVecSize == 0) { + launch_per_token_quant_fp8( + device.unwrap(), + static_cast(input.data_ptr()), + static_cast(output_q.data_ptr()), + static_cast(output_s.data_ptr()), + hidden_dim, + num_tokens); + } else { + launch_per_token_quant_fp8( + device.unwrap(), + static_cast(input.data_ptr()), + static_cast(output_q.data_ptr()), + static_cast(output_s.data_ptr()), + hidden_dim, + num_tokens); + } +} + +} // namespace sglang diff --git a/python/sglang/kernels/ops/quantization/__init__.py b/python/sglang/kernels/ops/quantization/__init__.py index 2b389d604..2032116e6 100644 --- a/python/sglang/kernels/ops/quantization/__init__.py +++ b/python/sglang/kernels/ops/quantization/__init__.py @@ -21,14 +21,15 @@ _CUDA = frozenset({CapabilityRequirement.CUDA}) register_kernel( KernelSpec( op="quantization.sgl_per_token_quant_fp8", - backend=KernelBackend.AOT, - target="sgl_kernel:sgl_per_token_quant_fp8", + backend=KernelBackend.JIT, + target="sglang.kernels.ops.quantization.per_token_quant_fp8:per_token_quant_fp8", + capabilities=_CUDA, format_signature=FormatSignature( supported_dtypes=("float8_e4m3fn",), in_place=True, description="per-token FP8 quantization into output_q/output_s", ), - description="Per-token FP8 quantization (sgl_kernel wheel).", + description="Per-token FP8 quantization (sglang.kernels.jit).", ) ) # fp8 / int8 are legacy aliases of the same 8bit kernel in the wheel; register @@ -78,7 +79,7 @@ def sgl_per_token_quant_fp8( output_s: torch.Tensor, ) -> None: """Per-token FP8 quantization, writing into ``output_q`` / ``output_s``.""" - return get_kernel("quantization.sgl_per_token_quant_fp8", KernelBackend.AOT)( + return get_kernel("quantization.sgl_per_token_quant_fp8", KernelBackend.JIT)( input, output_q, output_s ) diff --git a/python/sglang/kernels/ops/quantization/fp8_kernel.py b/python/sglang/kernels/ops/quantization/fp8_kernel.py index b0b8dec6b..6e3253fc4 100644 --- a/python/sglang/kernels/ops/quantization/fp8_kernel.py +++ b/python/sglang/kernels/ops/quantization/fp8_kernel.py @@ -46,7 +46,7 @@ _is_cpu = is_cpu() _is_musa = is_musa() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip -if _is_cuda or _is_musa: +if _is_cuda: from sglang.kernels.ops.quantization import ( per_token_group_quant, sgl_per_token_quant_fp8, @@ -55,6 +55,13 @@ if _is_cuda or _is_musa: per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8, ) +if _is_musa: + from sgl_kernel import sgl_per_token_quant_fp8 + + from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import ( + per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8, + ) + if _is_musa: # per_token_group_quant is CUDA-only JIT; MUSA keeps the AOT v2 group-quant op. from sglang.kernels.ops.quantization import sgl_per_token_group_quant_8bit @@ -2191,10 +2198,3 @@ def triton_scaled_mm( ) return result.to(out_dtype) - - -if _is_cuda: - - @register_fake_if_exists("sgl_kernel::sgl_per_token_quant_fp8") - def _(input, output_q, output_s): - return diff --git a/python/sglang/kernels/ops/quantization/per_token_quant_fp8.py b/python/sglang/kernels/ops/quantization/per_token_quant_fp8.py new file mode 100644 index 000000000..afb81a537 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/per_token_quant_fp8.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + get_jit_cuda_arch, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_per_token_quant_fp8_module(dtype: torch.dtype) -> Module: + if dtype not in (torch.float16, torch.bfloat16, torch.float32): + raise RuntimeError( + f"Unsupported dtype {dtype}. Supported: float16, bfloat16, float32" + ) + arch = get_jit_cuda_arch() + use_fast_math = (arch.major, arch.minor) == (9, 0) + math_mode = "fast_math" if use_fast_math else "precise_math" + args = make_cpp_args(dtype) + return load_jit( + "per_token_quant_fp8", + math_mode, + *args, + cuda_files=["gemm/per_token_quant_fp8.cuh"], + cuda_wrappers=[("per_token_quant_fp8", f"per_token_quant_fp8<{args}>")], + extra_cuda_cflags=["--use_fast_math"] if use_fast_math else [], + ) + + +@register_custom_op( + op_name="per_token_quant_fp8", + mutates_args=["output_q", "output_s"], +) +def per_token_quant_fp8( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, +) -> None: + """Dynamically quantize each row to FP8 E4M3.""" + module = _jit_per_token_quant_fp8_module(input.dtype) + module.per_token_quant_fp8(input, output_q, output_s.view(output_s.shape[0], 1)) diff --git a/test/registered/kernels/benchmark/quantization/bench_per_token_quant_fp8.py b/test/registered/kernels/benchmark/quantization/bench_per_token_quant_fp8.py new file mode 100644 index 000000000..b0d8ec079 --- /dev/null +++ b/test/registered/kernels/benchmark/quantization/bench_per_token_quant_fp8.py @@ -0,0 +1,42 @@ +import torch +from sgl_kernel import sgl_per_token_quant_fp8 as aot_per_token_quant_fp8 + +from sglang.kernels.jit.benchmark import marker +from sglang.kernels.jit.benchmark.utils import create_random +from sglang.kernels.ops.quantization.per_token_quant_fp8 import per_token_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + + +def _jit_quant(input, output, scale): + per_token_quant_fp8(input, output, scale) + + +FN_MAP = { + "jit": _jit_quant, + "aot": aot_per_token_quant_fp8, +} + + +@marker.parametrize("num_tokens", [1, 39, 128, 512, 1392, 7807], [39, 1392]) +@marker.parametrize("hidden_dim", [512, 1076, 1368, 1536, 2048, 4096], [1536]) +@marker.parametrize("dtype", [torch.float16, torch.bfloat16]) +@marker.benchmark("impl", ["jit", "aot"]) +def benchmark(num_tokens: int, hidden_dim: int, dtype: torch.dtype, impl: str): + input = create_random(num_tokens, hidden_dim, dtype=dtype) + output = torch.empty_like(input, dtype=torch.float8_e4m3fn) + scale = torch.empty((num_tokens, 1), dtype=torch.float32, device="cuda") + return marker.do_bench( + FN_MAP[impl], + input_args=(input, output, scale), + memory_args=(input,), + memory_output=(output, scale), + graph_clone_args=(0,), + ) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/kernels/ops/quantization/test_per_token_quant_fp8.py b/test/registered/kernels/ops/quantization/test_per_token_quant_fp8.py new file mode 100644 index 000000000..386c1d388 --- /dev/null +++ b/test/registered/kernels/ops/quantization/test_per_token_quant_fp8.py @@ -0,0 +1,120 @@ +import sys + +import pytest +import torch +from sgl_kernel import sgl_per_token_quant_fp8 as aot_per_token_quant_fp8 + +from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant +from sglang.kernels.ops.quantization.per_token_quant_fp8 import per_token_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=16, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=16, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +register_cuda_ci(est_time=30, stage="nightly", runner_config="1-gpu-large") + + +def _run_impl(input: torch.Tensor, *, use_jit: bool): + output = torch.empty_like(input, dtype=torch.float8_e4m3fn) + scale = torch.empty((input.shape[0], 1), dtype=torch.float32, device="cuda") + if use_jit: + per_token_quant_fp8(input, output, scale) + else: + aot_per_token_quant_fp8(input, output, scale) + return output, scale + + +def _assert_bitwise_equal(actual: torch.Tensor, expected: torch.Tensor): + assert torch.equal(actual.view(torch.uint8), expected.view(torch.uint8)) + + +def _warp_dispatch_num_tokens() -> int: + return torch.cuda.get_device_properties(0).multi_processor_count * 16 + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("dispatch", ["cta", "warp"]) +@pytest.mark.parametrize("hidden_dim", [1076, 1368]) +def test_per_token_quant_fp8_is_bit_exact(dtype, dispatch, hidden_dim): + """The JIT migration must preserve every output and scale bit from AOT.""" + num_tokens = 39 if dispatch == "cta" else _warp_dispatch_num_tokens() + input = torch.rand((num_tokens, hidden_dim), dtype=dtype, device="cuda") + + actual_output, actual_scale = _run_impl(input, use_jit=True) + expected_output, expected_scale = _run_impl(input, use_jit=False) + + _assert_bitwise_equal(actual_scale, expected_scale) + _assert_bitwise_equal(actual_output, expected_output) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("dispatch", ["cta", "warp"]) +def test_per_token_quant_fp8_zero_rows_are_bit_exact(dtype, dispatch): + """Zero-scale behavior differs by legacy dispatch and must remain unchanged.""" + num_tokens = 1 if dispatch == "cta" else _warp_dispatch_num_tokens() + input = torch.zeros((num_tokens, 512), dtype=dtype, device="cuda") + + actual_output, actual_scale = _run_impl(input, use_jit=True) + expected_output, expected_scale = _run_impl(input, use_jit=False) + + _assert_bitwise_equal(actual_scale, expected_scale) + _assert_bitwise_equal(actual_output, expected_output) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("dispatch", ["cta", "warp"]) +def test_per_token_quant_fp8_midpoints_are_bit_exact(dtype, dispatch): + """FP8 rounding ties must select the same representable value as AOT.""" + num_tokens = 1 if dispatch == "cta" else _warp_dispatch_num_tokens() + midpoint_values = torch.tensor( + [448.0, 1.0625, 1.1875, 1.375, -1.0625, -1.1875, -1.375], + dtype=dtype, + device="cuda", + ) + input = midpoint_values.repeat(num_tokens, 512 // midpoint_values.numel() + 1)[ + :, :512 + ].contiguous() + + actual_output, actual_scale = _run_impl(input, use_jit=True) + expected_output, expected_scale = _run_impl(input, use_jit=False) + + _assert_bitwise_equal(actual_scale, expected_scale) + _assert_bitwise_equal(actual_output, expected_output) + + +def test_scaled_fp8_quant_accepts_padded_outputs(): + """Dynamic per-token quantization supports the serving padding contract.""" + input = torch.rand((1, 512), dtype=torch.float16, device="cuda") + + output, scale = scaled_fp8_quant( + input, num_token_padding=17, use_per_token_if_dynamic=True + ) + expected_output, expected_scale = _run_impl(input, use_jit=False) + + assert output.shape == (17, 512) + assert scale.shape == (17, 1) + _assert_bitwise_equal(output[:1], expected_output) + _assert_bitwise_equal(scale[:1], expected_scale) + + +def test_per_token_quant_fp8_preserves_padded_tail(): + input = torch.rand((1, 512), dtype=torch.float16, device="cuda") + output = torch.full((17, 512), 1.0, dtype=torch.float8_e4m3fn, device="cuda") + scale = torch.full((17, 1), 2.0, dtype=torch.float32, device="cuda") + + per_token_quant_fp8(input, output, scale) + + assert torch.all(output[1:].float() == 1.0) + assert torch.all(scale[1:] == 2.0) + + +def test_per_token_quant_fp8_rejects_unsupported_dtype(): + input = torch.ones((1, 512), dtype=torch.int32, device="cuda") + output = torch.empty((1, 512), dtype=torch.float8_e4m3fn, device="cuda") + scale = torch.empty((1, 1), dtype=torch.float32, device="cuda") + + with pytest.raises(RuntimeError, match="Unsupported dtype"): + per_token_quant_fp8(input, output, scale) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"]))