diff --git a/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit_v2.cuh b/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit_v2.cuh new file mode 100644 index 000000000..cad7a9a92 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit_v2.cuh @@ -0,0 +1,531 @@ +// JIT port of the AOT sgl_per_token_group_quant_8bit_v2 (sgl-kernel). +// +// Same math as the AOT v2 kernel (256-bit vectorized loads, 8 threads/128-group, +// PDL, NaiveScheduler + MaskedLayoutScheduler, ue8m0/float scales, fp8/int8 +// output, fused silu+mul) so it is a drop-in replacement; the only changes vs the +// AOT source are the launcher (tvm::ffi::TensorView + TensorMatcher + the JIT +// LaunchKernel/PDL helpers) and the FP8 type alias. +#include // TensorMatcher, SymbolicSize/Device +#include // RuntimeCheck, Panic + +#include // LaunchKernel, fp8_e4m3_t, SGL_DEVICE, device::PDLWaitPrimary/TriggerSecondary +#include // device::warp::reduce_max + +#include + +#include +#include +#include + +namespace { + +constexpr float LOCAL_ABSMAX_ABS = 1e-10f; +constexpr uint32_t INPUT_PRIMARY_VEC_NUM_BYTES = 32; + +template +SGL_DEVICE float GroupReduceMax(float val) { + static_assert( + (THREADS_PER_SUBWARP & (THREADS_PER_SUBWARP - 1)) == 0 && THREADS_PER_SUBWARP <= 16 && THREADS_PER_SUBWARP >= 1, + "THREADS_PER_SUBWARP must be 1, 2, 4, 8, or 16"); + // Reduce within this thread's contiguous THREADS_PER_SUBWARP-lane subgroup via + // the shared warp primitive, but pass an explicit subgroup mask instead of its + // default 0xffffffff: the block can be < 32 lanes (subwarps_per_block * + // THREADS_PER_SUBWARP, e.g. 1..16), where a full-warp mask names non-existent + // lanes and is UB / can hang. + constexpr device::warp::mask_t kSub = (device::warp::mask_t{1} << THREADS_PER_SUBWARP) - 1; + const device::warp::mask_t mask = kSub << (THREADS_PER_SUBWARP * ((threadIdx.x % 32) / THREADS_PER_SUBWARP)); + return device::warp::reduce_max(val, mask); +} + +SGL_DEVICE float silu(const float& val) { + // Match the AOT v2 kernel: tanh-based silu on SM100+ (Blackwell), exp-based + // elsewhere, so the fused silu+mul output stays bit-identical to the AOT op. +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + float half = 0.5f * val; + float t = __tanhf(half); + return half * (1.0f + t); +#else + return val / (1.0f + __expf(-val)); +#endif +} + +SGL_DEVICE float2 fmul2_rn(float2 a, float2 b) { + // Match the AOT v2 kernel: use the __fmul2_rn intrinsic on SM100+. +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + return __fmul2_rn(a, b); +#else + float2 result; + result.x = a.x * b.x; + result.y = a.y * b.y; + return result; +#endif +} + +// Copied from DeepEP. +SGL_DEVICE float fast_pow2(int x) { + uint32_t bits_x = (x + 127) << 23; + return __uint_as_float(bits_x); // type-safe bit cast (no strict-aliasing UB) +} +SGL_DEVICE int fast_log2_ceil(float x) { + auto bits_x = __float_as_uint(x); // type-safe bit cast (no strict-aliasing UB) + auto exp_x = (bits_x >> 23) & 0xff; + auto man_bits = bits_x & ((1 << 23) - 1); + return exp_x - 127 + (man_bits != 0); +} + +template +struct DtypeInfo; +template <> +struct DtypeInfo { + static constexpr float MIN = -128; + static constexpr float MAX = 127; +}; +template <> +struct DtypeInfo { + static constexpr float MIN = -448; + static constexpr float MAX = 448; +}; + +template +SGL_DEVICE void calculate_fp8_scales(float amax, float& scale, float& scale_inv) { + constexpr float MAX_8BIT_INV = 1.0f / dtype_info::MAX; + if constexpr (ROUND_SCALE) { + auto exp_scale_inv = fast_log2_ceil(amax * MAX_8BIT_INV); + scale = fast_pow2(-exp_scale_inv); + scale_inv = fast_pow2(exp_scale_inv); + } else { + scale_inv = amax * MAX_8BIT_INV; + scale = dtype_info::MAX / amax; + } +} + +template > +SGL_DEVICE OUT_DTYPE_T extract_required_scale_format(float value) { + if constexpr (SCALE_UE8M0) { + return static_cast(__float_as_uint(value) >> 23); + } else { + return value; + } +} + +template +SGL_DEVICE int compute_input_group_start_offset( + int expert_idx, + int token_idx, + int hidden_dim_group_idx, + int hidden_size, + int num_tokens_per_expert, + int group_size) { + return expert_idx * num_tokens_per_expert * hidden_size * (FUSE_SILU_AND_MUL ? 2 : 1) + + token_idx * hidden_size * (FUSE_SILU_AND_MUL ? 2 : 1) + hidden_dim_group_idx * group_size; +} + +struct NaiveScheduler { + static void compute_exec_config( + int threads_per_subwarp, + int num_local_experts, + int hidden_dim_num_groups, + int num_groups, + int& subwarps_per_block, + dim3& grid, + dim3& block) { + subwarps_per_block = (num_groups % 16 == 0) ? 16 + : (num_groups % 8 == 0) ? 8 + : (num_groups % 4 == 0) ? 4 + : (num_groups % 2 == 0) ? 2 + : 1; + grid = dim3(num_groups / subwarps_per_block); + block = dim3(subwarps_per_block * threads_per_subwarp); + } + + template + SGL_DEVICE static void execute( + const int subwarps_per_block, + const int hidden_dim_num_groups, + const int32_t* masked_m, + const int num_tokens_per_expert, + FUNC fn) { + constexpr int expert_idx = 0; + const int64_t subwarp_id = threadIdx.x / THREADS_PER_SUBWARP; + const int lane_id = threadIdx.x % THREADS_PER_SUBWARP; + const int64_t group_id = static_cast(blockIdx.x) * subwarps_per_block + subwarp_id; + + int64_t input_group_start_offset; + if constexpr (!FUSE_SILU_AND_MUL) input_group_start_offset = group_id * GROUP_SIZE; + const int token_idx = group_id / hidden_dim_num_groups; + const int hidden_dim_group_idx = group_id % hidden_dim_num_groups; + if constexpr (FUSE_SILU_AND_MUL) { + const int hidden_size = hidden_dim_num_groups * GROUP_SIZE; + input_group_start_offset = compute_input_group_start_offset( + expert_idx, token_idx, hidden_dim_group_idx, hidden_size, num_tokens_per_expert, GROUP_SIZE); + } + fn(expert_idx, token_idx, hidden_dim_group_idx, lane_id, input_group_start_offset); + } +}; + +struct MaskedLayoutScheduler { + static constexpr int TOKEN_DIM_BLOCK_NUM_PER_EXPERT = 1024; + static constexpr int SUBWARPS_PER_BLOCK = 16; + + static void compute_exec_config( + int threads_per_subwarp, + int num_local_experts, + int hidden_dim_num_groups, + int num_groups, + int& subwarps_per_block, + dim3& grid, + dim3& block) { + subwarps_per_block = SUBWARPS_PER_BLOCK; + host::RuntimeCheck(hidden_dim_num_groups % subwarps_per_block == 0, "hidden_dim_num_groups not divisible by 16"); + grid = dim3(hidden_dim_num_groups / subwarps_per_block, TOKEN_DIM_BLOCK_NUM_PER_EXPERT, num_local_experts); + block = dim3(subwarps_per_block * threads_per_subwarp); + } + + template + SGL_DEVICE static void execute( + const int subwarps_per_block, + const int hidden_dim_num_groups, + const int32_t* masked_m, + const int num_tokens_per_expert, + FUNC fn) { + const int64_t subwarp_id = threadIdx.x / THREADS_PER_SUBWARP; + const int lane_id = threadIdx.x % THREADS_PER_SUBWARP; + const int expert_idx = blockIdx.z; + const int token_idx_start = blockIdx.y; + const int64_t hidden_dim_group_idx = static_cast(blockIdx.x) * SUBWARPS_PER_BLOCK + subwarp_id; + const int curr_expert_token_num = masked_m[expert_idx]; + for (int token_idx = token_idx_start; token_idx < curr_expert_token_num; + token_idx += TOKEN_DIM_BLOCK_NUM_PER_EXPERT) { + const int hidden_size = hidden_dim_num_groups * GROUP_SIZE; + const int64_t input_group_start_offset = compute_input_group_start_offset( + expert_idx, token_idx, hidden_dim_group_idx, hidden_size, num_tokens_per_expert, GROUP_SIZE); + fn(expert_idx, token_idx, hidden_dim_group_idx, lane_id, input_group_start_offset); + } + } +}; + +template < + typename SCHEDULER, + int GROUP_SIZE, + int THREADS_PER_SUBWARP, + typename T, + typename DST_DTYPE, + bool IS_COLUMN_MAJOR, + bool SCALE_UE8M0, + bool FUSE_SILU_AND_MUL, + bool kUsePDL, + typename scale_packed_t = std::conditional_t> +__global__ void per_token_group_quant_8bit_v2_kernel( + const T* __restrict__ input, + DST_DTYPE* __restrict__ output_q, + scale_packed_t* __restrict__ output_s, + const int32_t* __restrict__ masked_m, + const int subwarps_per_block, + const int hidden_dim_num_groups, + const int scale_expert_stride, + const int scale_hidden_stride, + const int num_tokens_per_expert) { + using dst_dtype_info = DtypeInfo; + using scale_element_t = std::conditional_t; + static_assert(sizeof(scale_packed_t) % sizeof(scale_element_t) == 0); + + device::PDLWaitPrimary(); + + SCHEDULER::template execute( + subwarps_per_block, + hidden_dim_num_groups, + masked_m, + num_tokens_per_expert, + [&](const int expert_idx, + const int token_idx, + const int hidden_dim_group_idx, + const int lane_id, + const int input_group_start_offset) { + constexpr uint32_t INPUT_PRIMARY_VEC_SIZE = INPUT_PRIMARY_VEC_NUM_BYTES / sizeof(T); + constexpr uint32_t INPUT_PRIMARY_INT4_SIZE = INPUT_PRIMARY_VEC_NUM_BYTES / sizeof(int4); + + const int offset_num_groups = expert_idx * num_tokens_per_expert * hidden_dim_num_groups + + token_idx * hidden_dim_num_groups + hidden_dim_group_idx; + + int4 input_primary_int4[INPUT_PRIMARY_INT4_SIZE]; + T* input_primary_vec = reinterpret_cast(input_primary_int4); + int4 input_secondary_int4[INPUT_PRIMARY_INT4_SIZE]; + T* input_secondary_vec = reinterpret_cast(input_secondary_int4); + +#pragma unroll + for (uint32_t j = 0; j < INPUT_PRIMARY_INT4_SIZE; ++j) { + // Ordinary 128-bit vectorized load (LDG.128); .nc gave no measurable + // gain on this streaming read-once kernel, so no inline asm. + input_primary_int4[j] = + reinterpret_cast(input + input_group_start_offset + lane_id * INPUT_PRIMARY_VEC_SIZE)[j]; + } + if constexpr (FUSE_SILU_AND_MUL) { + const int secondary_offset = hidden_dim_num_groups * GROUP_SIZE; +#pragma unroll + for (uint32_t j = 0; j < INPUT_PRIMARY_INT4_SIZE; ++j) { + input_secondary_int4[j] = reinterpret_cast( + input + input_group_start_offset + lane_id * INPUT_PRIMARY_VEC_SIZE + secondary_offset)[j]; + } + } + + constexpr int num_elems_per_pack = static_cast(sizeof(scale_packed_t) / sizeof(scale_element_t)); + scale_element_t* scale_output; + if constexpr (IS_COLUMN_MAJOR) { + constexpr int scale_token_stride = 1; + const int hidden_idx_packed = hidden_dim_group_idx / num_elems_per_pack; + const int pack_idx = hidden_dim_group_idx % num_elems_per_pack; + scale_output = reinterpret_cast(output_s) + + (expert_idx * scale_expert_stride * num_elems_per_pack + + hidden_idx_packed * scale_hidden_stride * num_elems_per_pack + + token_idx * scale_token_stride * num_elems_per_pack + pack_idx); + } else { + static_assert(!SCALE_UE8M0); + scale_output = output_s + offset_num_groups; + } + + if constexpr (IS_COLUMN_MAJOR and SCALE_UE8M0) { + const int remainder_num_groups = hidden_dim_num_groups % num_elems_per_pack; + if ((remainder_num_groups != 0) and (hidden_dim_group_idx == hidden_dim_num_groups - 1) and + (lane_id < num_elems_per_pack - remainder_num_groups)) { + const int shift = 1 + lane_id; + *(scale_output + shift) = 0; + } + } + + float local_absmax = LOCAL_ABSMAX_ABS; +#pragma unroll + for (uint32_t j = 0; j < INPUT_PRIMARY_VEC_SIZE; ++j) { + float val; + if constexpr (FUSE_SILU_AND_MUL) { + T val_lowprec = static_cast(silu(static_cast(input_primary_vec[j]))) * input_secondary_vec[j]; + val = static_cast(val_lowprec); + input_primary_vec[j] = val_lowprec; + } else { + val = static_cast(input_primary_vec[j]); + } + local_absmax = fmaxf(local_absmax, fabsf(val)); + } + + local_absmax = GroupReduceMax(local_absmax); + + float y_scale, y_scale_inv; + calculate_fp8_scales(local_absmax, y_scale, y_scale_inv); + float2 y_scale_repeated = {y_scale, y_scale}; + + if (lane_id == 0) { + *scale_output = extract_required_scale_format(y_scale_inv); + } + + int4 output_buf; + if constexpr (std::is_same_v) { + const auto output_buf_ptr = reinterpret_cast<__nv_fp8x2_storage_t*>(&output_buf); +#pragma unroll + for (uint32_t j = 0; j < INPUT_PRIMARY_VEC_SIZE; j += 2) { + float2 inputx2 = {static_cast(input_primary_vec[j]), static_cast(input_primary_vec[j + 1])}; + float2 outputx2 = fmul2_rn(inputx2, y_scale_repeated); + outputx2.x = fminf(fmaxf(outputx2.x, dst_dtype_info::MIN), dst_dtype_info::MAX); + outputx2.y = fminf(fmaxf(outputx2.y, dst_dtype_info::MIN), dst_dtype_info::MAX); + output_buf_ptr[j / 2] = __nv_cvt_float2_to_fp8x2(outputx2, __NV_SATFINITE, __NV_E4M3); + } + } else { + const auto output_buf_ptr = reinterpret_cast(&output_buf); +#pragma unroll + for (uint32_t j = 0; j < INPUT_PRIMARY_VEC_SIZE; ++j) { + float val = static_cast(input_primary_vec[j]); + float q_val = fminf(fmaxf(val * y_scale, dst_dtype_info::MIN), dst_dtype_info::MAX); + output_buf_ptr[j] = DST_DTYPE(q_val); + } + } + + // Ordinary 128-bit vectorized store (STG.128); no inline asm. + *reinterpret_cast(output_q + offset_num_groups * GROUP_SIZE + lane_id * INPUT_PRIMARY_VEC_SIZE) = + output_buf; + }); + + device::PDLTriggerSecondary(); +} + +// ---------------------------------------------------------------------------- +// Launcher (JIT). All shape-derived scalars are computed in the Python wrapper +// and passed in, so the C++ side only needs TensorView::data_ptr()/device(). +// Runtime combos (column_major / ue8m0 / silu / masked / group_size) are +// dispatched to the templated kernel; launch is PDL-aware via LaunchKernel. +// ---------------------------------------------------------------------------- +template +struct TypeTag { + using type = S; +}; + +template +struct PerTokenGroupQuant8bitV2Kernel { + template < + typename SCHEDULER, + int GROUP_SIZE, + int THREADS_PER_SUBWARP, + bool IS_COLUMN_MAJOR, + bool SCALE_UE8M0, + bool FUSE_SILU_AND_MUL> + static void launch( + const DLDevice& device, + dim3 grid, + dim3 block, + int subwarps_per_block, + int hidden_dim_num_groups, + int scale_expert_stride, + int scale_hidden_stride, + int num_tokens_per_expert, + const void* input, + void* output_q, + void* output_s, + const int32_t* masked_m) { + using scale_packed_t = std::conditional_t; + auto kernel = per_token_group_quant_8bit_v2_kernel< + SCHEDULER, + GROUP_SIZE, + THREADS_PER_SUBWARP, + T, + DST_DTYPE, + IS_COLUMN_MAJOR, + SCALE_UE8M0, + FUSE_SILU_AND_MUL, + kUsePDL>; + host::LaunchKernel(grid, block, device) + .enable_pdl(kUsePDL)( + kernel, + static_cast(input), + static_cast(output_q), + static_cast(output_s), + masked_m, + subwarps_per_block, + hidden_dim_num_groups, + scale_expert_stride, + scale_hidden_stride, + num_tokens_per_expert); + } + + template + static void dispatch_bools( + const DLDevice& device, + bool is_column_major, + bool scale_ue8m0, + bool fuse_silu_and_mul, + bool masked_layout, + int num_local_experts, + int hidden_dim_num_groups, + int num_groups, + int scale_expert_stride, + int scale_hidden_stride, + int num_tokens_per_expert, + const void* input, + void* output_q, + void* output_s, + const int32_t* masked_m) { + constexpr int THREADS_PER_SUBWARP = GROUP_SIZE / 16; + + auto launch_with_config = [&](auto sched_tag, auto colmajor_tag, auto ue8m0_tag, auto silu_tag) { + using SCHEDULER = typename decltype(sched_tag)::type; + int subwarps_per_block; + dim3 grid, block; + SCHEDULER::compute_exec_config( + THREADS_PER_SUBWARP, num_local_experts, hidden_dim_num_groups, num_groups, subwarps_per_block, grid, block); + launch< + SCHEDULER, + GROUP_SIZE, + THREADS_PER_SUBWARP, + decltype(colmajor_tag)::value, + decltype(ue8m0_tag)::value, + decltype(silu_tag)::value>( + device, + grid, + block, + subwarps_per_block, + hidden_dim_num_groups, + scale_expert_stride, + scale_hidden_stride, + num_tokens_per_expert, + input, + output_q, + output_s, + masked_m); + }; + + if (is_column_major) { + if (scale_ue8m0) { + if (fuse_silu_and_mul) { + if (masked_layout) + launch_with_config(TypeTag{}, std::true_type{}, std::true_type{}, std::true_type{}); + else + launch_with_config(TypeTag{}, std::true_type{}, std::true_type{}, std::true_type{}); + } else { + launch_with_config(TypeTag{}, std::true_type{}, std::true_type{}, std::false_type{}); + } + } else { + launch_with_config(TypeTag{}, std::true_type{}, std::false_type{}, std::false_type{}); + } + } else { + launch_with_config(TypeTag{}, std::false_type{}, std::false_type{}, std::false_type{}); + } + } + + static void + run(tvm::ffi::TensorView input, + tvm::ffi::TensorView output_q, + tvm::ffi::TensorView output_s, + tvm::ffi::TensorView masked_m, + int64_t group_size, + bool scale_ue8m0, + bool fuse_silu_and_mul, + bool masked_layout, + int64_t num_groups, + int64_t num_local_experts, + bool is_column_major, + int64_t hidden_dim_num_groups, + int64_t num_tokens_per_expert, + int64_t scale_expert_stride, + int64_t scale_hidden_stride) { + const DLDevice dev = input.device(); + const void* in = input.data_ptr(); + void* oq = output_q.data_ptr(); + void* os = output_s.data_ptr(); + const int32_t* masked_ptr = masked_layout ? static_cast(masked_m.data_ptr()) : nullptr; + + auto dispatch_gs = [&](auto gs_tag) { + constexpr int GS = decltype(gs_tag)::value; + static_assert((GS / 16) * INPUT_PRIMARY_VEC_NUM_BYTES == GS * static_cast(sizeof(T))); + dispatch_bools( + dev, + is_column_major, + scale_ue8m0, + fuse_silu_and_mul, + masked_layout, + static_cast(num_local_experts), + static_cast(hidden_dim_num_groups), + static_cast(num_groups), + static_cast(scale_expert_stride), + static_cast(scale_hidden_stride), + static_cast(num_tokens_per_expert), + in, + oq, + os, + masked_ptr); + }; + switch (group_size) { + case 16: + dispatch_gs(std::integral_constant{}); + break; + case 32: + dispatch_gs(std::integral_constant{}); + break; + case 64: + dispatch_gs(std::integral_constant{}); + break; + case 128: + dispatch_gs(std::integral_constant{}); + break; + default: + host::Panic("Unsupported group_size ", group_size); + } + } +}; + +} // namespace diff --git a/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py b/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py new file mode 100644 index 000000000..8c0ebbfc2 --- /dev/null +++ b/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.jit_kernel.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.kernel_api_logging import debug_kernel_api +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_module(in_dtype: torch.dtype, out_dtype: torch.dtype, use_pdl: bool) -> Module: + args = make_cpp_args(in_dtype, out_dtype, use_pdl) + return load_jit( + "per_token_group_quant_8bit_v2", + *args, + cuda_files=["gemm/per_token_group_quant_8bit_v2.cuh"], + cuda_wrappers=[ + ( + "per_token_group_quant_8bit_v2", + f"PerTokenGroupQuant8bitV2Kernel<{args}>::run", + ) + ], + # Match the AOT sgl-kernel build (-use_fast_math) so the FP8 scale + # division/rounding is bit-identical to sgl_per_token_group_quant_8bit_v2. + extra_cuda_cflags=["--use_fast_math"], + ) + + +@register_custom_op( + op_name="per_token_group_quant_8bit_v2", + mutates_args=["output_q", "output_s"], +) +def _per_token_group_quant_8bit_v2_custom_op( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + eps: float, + min_8bit: float, + max_8bit: float, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, +) -> None: + """Opaque custom-op boundary around the JIT v2 kernel. + + Registering this as a custom op (instead of calling the tvm-ffi module + directly) keeps torch.compile / piecewise-CUDA-graph from tracing into the + tvm-ffi ``Function.__call__`` (which Dynamo cannot trace). All shape-derived + scalars are computed here and passed to the kernel. + + Layouts (matching the AOT v2): + vanilla: input (num_tokens, hidden), output_q (num_tokens, hidden) + fuse_silu_and_mul: input (num_tokens, hidden*2), output_q (num_tokens, hidden) + fuse_silu_and_mul+masked: input (num_experts, tokens_pad, hidden*2), + output_q (num_experts, tokens_pad, hidden), masked_m (num_experts,) + """ + masked_layout = masked_m is not None + numel = input.numel() + num_groups = numel // group_size // (2 if fuse_silu_and_mul else 1) + if num_groups == 0: # empty input -> grid 0 -> cudaErrorInvalidConfiguration + return + num_local_experts = input.shape[0] if masked_layout else 1 + last = output_q.dim() - 1 + is_column_major = output_s.stride(last - 1) < output_s.stride(last) + hidden_dim_num_groups = output_q.shape[last] // group_size + num_tokens_per_expert = output_q.shape[last - 1] + scale_expert_stride = output_s.stride(0) if masked_layout else 0 + scale_hidden_stride = output_s.stride(last) + + module = _jit_module(input.dtype, output_q.dtype, is_arch_support_pdl()) + module.per_token_group_quant_8bit_v2( + input, + output_q, + output_s, + masked_m if masked_layout else input, # unused (nullptr) when not masked + int(group_size), + bool(scale_ue8m0), + bool(fuse_silu_and_mul), + bool(masked_layout), + int(num_groups), + int(num_local_experts), + bool(is_column_major), + int(hidden_dim_num_groups), + int(num_tokens_per_expert), + int(scale_expert_stride), + int(scale_hidden_stride), + ) + + +@debug_kernel_api +def per_token_group_quant_8bit_v2( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + eps: float, + min_8bit: float, + max_8bit: float, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, +) -> None: + """JIT port of sgl_per_token_group_quant_8bit_v2 (full feature parity). + + Wraps the registered custom op so torch.compile / piecewise CUDA graph treat + the tvm-ffi kernel call as an opaque boundary. + """ + _per_token_group_quant_8bit_v2_custom_op( + input=input, + output_q=output_q, + output_s=output_s, + group_size=group_size, + eps=eps, + min_8bit=min_8bit, + max_8bit=max_8bit, + scale_ue8m0=scale_ue8m0, + fuse_silu_and_mul=fuse_silu_and_mul, + masked_m=masked_m, + ) diff --git a/python/sglang/srt/layers/quantization/fp8_kernel.py b/python/sglang/srt/layers/quantization/fp8_kernel.py index 27c2c63a7..2450697f9 100644 --- a/python/sglang/srt/layers/quantization/fp8_kernel.py +++ b/python/sglang/srt/layers/quantization/fp8_kernel.py @@ -73,6 +73,9 @@ if _is_cuda or _is_musa: from sglang.jit_kernel.per_token_group_quant_8bit import ( per_token_group_quant_8bit as sgl_per_token_group_quant_8bit_jit, ) + from sglang.jit_kernel.per_token_group_quant_8bit_v2 import ( + per_token_group_quant_8bit_v2 as sgl_per_token_group_quant_8bit_jit_v2, + ) if _is_hip: _has_vllm = False @@ -531,7 +534,10 @@ def sglang_per_token_group_quant_fp8( if x.shape[0] > 0: # Temporary if enable_sgl_per_token_group_quant_8bit: - if enable_v2: + if enable_v2 and _is_musa: + # The JIT v2 .cuh uses CUDA-only inline PTX (ld/st.global.v4) and + # has no MUSA fallback, so keep MUSA on the AOT v2 op, which + # carries the USE_MUSA vector load/store fallbacks. sgl_per_token_group_quant_8bit( x, x_q, @@ -545,6 +551,19 @@ def sglang_per_token_group_quant_fp8( masked_m, enable_v2=True, ) + elif enable_v2: + sgl_per_token_group_quant_8bit_jit_v2( + x, + x_q, + x_s, + group_size, + eps, + fp8_min, + fp8_max, + scale_ue8m0=scale_ue8m0, + fuse_silu_and_mul=fuse_silu_and_mul, + masked_m=masked_m, + ) else: sgl_per_token_group_quant_8bit_jit( input=x, diff --git a/test/registered/jit/benchmark/bench_per_token_group_quant_8bit_v2.py b/test/registered/jit/benchmark/bench_per_token_group_quant_8bit_v2.py new file mode 100644 index 000000000..983c5b697 --- /dev/null +++ b/test/registered/jit/benchmark/bench_per_token_group_quant_8bit_v2.py @@ -0,0 +1,65 @@ +import torch +from sgl_kernel import sgl_per_token_group_quant_8bit + +from sglang.jit_kernel.benchmark import marker +from sglang.jit_kernel.benchmark.utils import create_random +from sglang.jit_kernel.per_token_group_quant_8bit_v2 import ( + per_token_group_quant_8bit_v2, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + create_per_token_group_quant_fp8_output_scale, + fp8_dtype, + fp8_max, + fp8_min, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=6, suite="base-b-kernel-benchmark-1-gpu-large") + +G = 128 +HIDDEN = 4096 + + +def _aot_v2(x, x_q, x_s): + # Low-level AOT op writing into the same preallocated x_q/x_s as the JIT + # path, so this is a kernel-vs-kernel comparison (no wrapper / no realloc). + sgl_per_token_group_quant_8bit( + x, + x_q, + x_s, + G, + 1e-10, + float(fp8_min), + float(fp8_max), + False, # scale_ue8m0 + False, # fuse_silu_and_mul + None, # masked_m + enable_v2=True, + ) + + +def _jit_v2(x, x_q, x_s): + per_token_group_quant_8bit_v2(x, x_q, x_s, G, 1e-10, float(fp8_min), float(fp8_max)) + + +FN = {"aot_v2": _aot_v2, "jit_v2": _jit_v2} + + +@marker.parametrize("num_tokens", [1, 8, 64, 512, 4096], ci_vals=[1, 512]) +@marker.benchmark("impl", ["aot_v2", "jit_v2"]) +def benchmark(num_tokens: int, impl: str): + x = create_random(num_tokens, HIDDEN) + x_q = torch.empty(num_tokens, HIDDEN, device="cuda", dtype=fp8_dtype) + x_s = create_per_token_group_quant_fp8_output_scale( + x_shape=(num_tokens, HIDDEN), + device="cuda", + group_size=G, + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=False, + ) + return marker.do_bench(FN[impl], input_args=(x, x_q, x_s), graph_clone_args=(0,)) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/jit/test_per_token_group_quant_8bit_v2.py b/test/registered/jit/test_per_token_group_quant_8bit_v2.py new file mode 100644 index 000000000..7d5110c7f --- /dev/null +++ b/test/registered/jit/test_per_token_group_quant_8bit_v2.py @@ -0,0 +1,141 @@ +import pytest +import torch +from sgl_kernel import sgl_per_token_group_quant_8bit # AOT v2 reference op + +from sglang.jit_kernel.per_token_group_quant_8bit_v2 import ( + per_token_group_quant_8bit_v2, +) +from sglang.srt.layers.quantization.fp8_kernel import ( + create_per_token_group_quant_fp8_output_scale, + fp8_dtype, + fp8_max, + fp8_min, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=90, suite="base-b-kernel-unit-1-gpu-large") + +G = 128 + + +def _alloc(x_shape, scale_ue8m0): + """Pre-allocated (zeroed) output_q + output_s for a given input/output shape. + Zeroing makes the unwritten (padding / aligned) regions compare equal.""" + x_q = torch.zeros(x_shape, device="cuda", dtype=fp8_dtype) + x_s = create_per_token_group_quant_fp8_output_scale( + x_shape=x_shape, + device="cuda", + group_size=G, + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=scale_ue8m0, + ) + x_s.zero_() + return x_q, x_s + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("num_tokens", [1, 7, 64, 333]) +@pytest.mark.parametrize("hidden", [128, 2048, 4096]) +@pytest.mark.parametrize("fuse_silu_and_mul", [False, True]) +@pytest.mark.parametrize("scale_ue8m0", [False, True]) +def test_v2_jit_matches_aot(dtype, num_tokens, hidden, fuse_silu_and_mul, scale_ue8m0): + """JIT v2 must be bit-exact with the AOT v2 across vanilla/silu and float/ue8m0 + scales (NaiveScheduler).""" + torch.manual_seed( + hidden + num_tokens + int(fuse_silu_and_mul) + 7 * int(scale_ue8m0) + ) + in_hidden = hidden * (2 if fuse_silu_and_mul else 1) + x = torch.randn(num_tokens, in_hidden, device="cuda", dtype=dtype) + out_shape = (num_tokens, hidden) + + q_ref, s_ref = _alloc(out_shape, scale_ue8m0) + sgl_per_token_group_quant_8bit( + x, + q_ref, + s_ref, + G, + 1e-10, + float(fp8_min), + float(fp8_max), + scale_ue8m0, + fuse_silu_and_mul, + None, + enable_v2=True, + ) + + x_q, x_s = _alloc(out_shape, scale_ue8m0) + per_token_group_quant_8bit_v2( + x, + x_q, + x_s, + G, + 1e-10, + float(fp8_min), + float(fp8_max), + scale_ue8m0=scale_ue8m0, + fuse_silu_and_mul=fuse_silu_and_mul, + ) + torch.cuda.synchronize() + + assert torch.equal(x_q.view(torch.int8), q_ref.view(torch.int8)), "fp8 codes differ" + assert torch.equal(x_s, s_ref), "scales differ" + + +# Masked (EP-MoE) path: the v2 op only has a masked scheduler for the +# column-major + ue8m0 + fused-silu+mul + masked combination. Input is 3D +# [num_experts, tokens_padded, hidden*2]; only tokens < masked_m[e] are processed +# (padding left untouched → zeros in both). Compare JIT vs AOT bit-exact. +@pytest.mark.parametrize("num_experts", [2, 5]) +@pytest.mark.parametrize("hidden", [2048, 4096]) +@pytest.mark.parametrize("tokens_pad", [128, 384]) +def test_v2_jit_masked_matches_aot(num_experts, hidden, tokens_pad): + torch.manual_seed(num_experts * 1000 + hidden + tokens_pad) + x = torch.randn( + num_experts, tokens_pad, hidden * 2, device="cuda", dtype=torch.bfloat16 + ) + masked_m = torch.randint( + 0, tokens_pad + 1, (num_experts,), device="cuda", dtype=torch.int32 + ) + out_shape = (num_experts, tokens_pad, hidden) + + q_ref, s_ref = _alloc(out_shape, scale_ue8m0=True) + sgl_per_token_group_quant_8bit( + x, + q_ref, + s_ref, + G, + 1e-10, + float(fp8_min), + float(fp8_max), + True, + True, + masked_m, + enable_v2=True, + ) + + x_q, x_s = _alloc(out_shape, scale_ue8m0=True) + per_token_group_quant_8bit_v2( + x, + x_q, + x_s, + G, + 1e-10, + float(fp8_min), + float(fp8_max), + scale_ue8m0=True, + fuse_silu_and_mul=True, + masked_m=masked_m, + ) + torch.cuda.synchronize() + + assert torch.equal( + x_q.view(torch.int8), q_ref.view(torch.int8) + ), "masked fp8 differ" + assert torch.equal(x_s, s_ref), "masked scales differ" + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"]))