[jit-kernel] Support per token group quant 8bit v2 jit kernel (#27449)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-06-13 12:15:09 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent 053e153183
commit eb18416f9f
5 changed files with 887 additions and 1 deletions
@@ -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 <sgl_kernel/tensor.h> // TensorMatcher, SymbolicSize/Device
#include <sgl_kernel/utils.h> // RuntimeCheck, Panic
#include <sgl_kernel/utils.cuh> // LaunchKernel, fp8_e4m3_t, SGL_DEVICE, device::PDLWaitPrimary/TriggerSecondary
#include <sgl_kernel/warp.cuh> // device::warp::reduce_max
#include <tvm/ffi/container/tensor.h>
#include <cstdint>
#include <cuda_fp8.h>
#include <type_traits>
namespace {
constexpr float LOCAL_ABSMAX_ABS = 1e-10f;
constexpr uint32_t INPUT_PRIMARY_VEC_NUM_BYTES = 32;
template <int THREADS_PER_SUBWARP>
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<THREADS_PER_SUBWARP>(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 <typename T>
struct DtypeInfo;
template <>
struct DtypeInfo<int8_t> {
static constexpr float MIN = -128;
static constexpr float MAX = 127;
};
template <>
struct DtypeInfo<fp8_e4m3_t> {
static constexpr float MIN = -448;
static constexpr float MAX = 448;
};
template <bool ROUND_SCALE, typename dtype_info>
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 <bool SCALE_UE8M0, typename OUT_DTYPE_T = std::conditional_t<SCALE_UE8M0, uint8_t, float>>
SGL_DEVICE OUT_DTYPE_T extract_required_scale_format(float value) {
if constexpr (SCALE_UE8M0) {
return static_cast<uint8_t>(__float_as_uint(value) >> 23);
} else {
return value;
}
}
template <bool FUSE_SILU_AND_MUL>
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 <bool FUSE_SILU_AND_MUL, int GROUP_SIZE, int THREADS_PER_SUBWARP, typename FUNC>
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<int64_t>(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<FUSE_SILU_AND_MUL>(
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 <bool FUSE_SILU_AND_MUL, int GROUP_SIZE, int THREADS_PER_SUBWARP, typename FUNC>
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<int64_t>(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<FUSE_SILU_AND_MUL>(
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<SCALE_UE8M0, uint32_t, float>>
__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<DST_DTYPE>;
using scale_element_t = std::conditional_t<SCALE_UE8M0, uint8_t, float>;
static_assert(sizeof(scale_packed_t) % sizeof(scale_element_t) == 0);
device::PDLWaitPrimary<kUsePDL>();
SCHEDULER::template execute<FUSE_SILU_AND_MUL, GROUP_SIZE, THREADS_PER_SUBWARP>(
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<T*>(input_primary_int4);
int4 input_secondary_int4[INPUT_PRIMARY_INT4_SIZE];
T* input_secondary_vec = reinterpret_cast<T*>(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<const int4*>(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<const int4*>(
input + input_group_start_offset + lane_id * INPUT_PRIMARY_VEC_SIZE + secondary_offset)[j];
}
}
constexpr int num_elems_per_pack = static_cast<int>(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<scale_element_t*>(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<T>(silu(static_cast<float>(input_primary_vec[j]))) * input_secondary_vec[j];
val = static_cast<float>(val_lowprec);
input_primary_vec[j] = val_lowprec;
} else {
val = static_cast<float>(input_primary_vec[j]);
}
local_absmax = fmaxf(local_absmax, fabsf(val));
}
local_absmax = GroupReduceMax<THREADS_PER_SUBWARP>(local_absmax);
float y_scale, y_scale_inv;
calculate_fp8_scales<SCALE_UE8M0, dst_dtype_info>(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<SCALE_UE8M0>(y_scale_inv);
}
int4 output_buf;
if constexpr (std::is_same_v<DST_DTYPE, fp8_e4m3_t>) {
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<float>(input_primary_vec[j]), static_cast<float>(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<DST_DTYPE*>(&output_buf);
#pragma unroll
for (uint32_t j = 0; j < INPUT_PRIMARY_VEC_SIZE; ++j) {
float val = static_cast<float>(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<int4*>(output_q + offset_num_groups * GROUP_SIZE + lane_id * INPUT_PRIMARY_VEC_SIZE) =
output_buf;
});
device::PDLTriggerSecondary<kUsePDL>();
}
// ----------------------------------------------------------------------------
// 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 <typename S>
struct TypeTag {
using type = S;
};
template <typename T, typename DST_DTYPE, bool kUsePDL>
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<SCALE_UE8M0, uint32_t, float>;
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<const T*>(input),
static_cast<DST_DTYPE*>(output_q),
static_cast<scale_packed_t*>(output_s),
masked_m,
subwarps_per_block,
hidden_dim_num_groups,
scale_expert_stride,
scale_hidden_stride,
num_tokens_per_expert);
}
template <int GROUP_SIZE>
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<MaskedLayoutScheduler>{}, std::true_type{}, std::true_type{}, std::true_type{});
else
launch_with_config(TypeTag<NaiveScheduler>{}, std::true_type{}, std::true_type{}, std::true_type{});
} else {
launch_with_config(TypeTag<NaiveScheduler>{}, std::true_type{}, std::true_type{}, std::false_type{});
}
} else {
launch_with_config(TypeTag<NaiveScheduler>{}, std::true_type{}, std::false_type{}, std::false_type{});
}
} else {
launch_with_config(TypeTag<NaiveScheduler>{}, 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<const int32_t*>(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<int>(sizeof(T)));
dispatch_bools<GS>(
dev,
is_column_major,
scale_ue8m0,
fuse_silu_and_mul,
masked_layout,
static_cast<int>(num_local_experts),
static_cast<int>(hidden_dim_num_groups),
static_cast<int>(num_groups),
static_cast<int>(scale_expert_stride),
static_cast<int>(scale_hidden_stride),
static_cast<int>(num_tokens_per_expert),
in,
oq,
os,
masked_ptr);
};
switch (group_size) {
case 16:
dispatch_gs(std::integral_constant<int, 16>{});
break;
case 32:
dispatch_gs(std::integral_constant<int, 32>{});
break;
case 64:
dispatch_gs(std::integral_constant<int, 64>{});
break;
case 128:
dispatch_gs(std::integral_constant<int, 128>{});
break;
default:
host::Panic("Unsupported group_size ", group_size);
}
}
};
} // namespace
@@ -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,
)
@@ -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,
@@ -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()
@@ -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"]))