[Fix/Kernel] Add JIT rmsnorm_hf kernel to fix transformers backend MMLU accuracy regression (#22931)

Co-authored-by: SGLang CI <ci@sglang.ai>
This commit is contained in:
Jimmy Shong
2026-04-23 12:00:31 +08:00
committed by GitHub
co-authored by SGLang CI
parent 267c2c0849
commit 68a8ed9b11
5 changed files with 510 additions and 0 deletions
@@ -0,0 +1,253 @@
/**
* RMSNorm with HuggingFace semantics:
* out[i] = weight[i] * cast_dtype( rsqrt(mean_j(x[j]^2) + eps) * x[i] )
*
* vs. standard rmsnorm: the normalized x is rounded to the activation dtype
* BEFORE the weight multiply (not after). The multiply itself is done in fp32
* either way; the load-bearing step is the intermediate rounding. Required
* for HF `LlamaRMSNorm` parity under weight-only quantization.
*
* Two launch configs:
* - Warp kernel: 32 threads/row for small hidden sizes (q/k norms).
* - CTA kernel: 512-thread scalar-strided with register cache (token norms).
*/
#include <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
#include <sgl_kernel/utils.h> // For RuntimeCheck
#include <sgl_kernel/math.cuh> // For device::math::rsqrt
#include <sgl_kernel/runtime.cuh> // For runtime::get_blocks_per_sm, get_sm_count
#include <sgl_kernel/utils.cuh> // For LaunchKernel, SGL_DEVICE, type aliases, PDL, cast
#include <sgl_kernel/warp.cuh> // For warp::reduce_sum
#include <tvm/ffi/container/tensor.h>
namespace {
struct RMSNormHFParams {
const void* input;
const void* __restrict__ weight;
void* output;
int64_t input_stride;
int64_t output_stride;
uint32_t num_tokens;
float eps;
};
// ---------------------------------------------------------------------------
// Warp kernel: one warp per row, for small hidden sizes (e.g. q/k norms at
// head_dim ∈ {32, 64, 96, 128, 256}). No shared memory, no block reduce —
// warp reduce is sufficient. Grid-strided over rows.
// ---------------------------------------------------------------------------
template <int64_t kDim, bool kUsePDL, typename Float>
__global__ __launch_bounds__(32) void rmsnorm_hf_warp_kernel(const RMSNormHFParams __grid_constant__ params) {
using namespace device;
constexpr int kElemsPerThread = kDim / kWarpThreads;
const auto& [input, weight_ptr, output, input_stride, output_stride, num_tokens, eps] = params;
const auto wr = static_cast<const Float*>(weight_ptr);
PDLWaitPrimary<kUsePDL>();
for (uint32_t row = blockIdx.x; row < num_tokens; row += gridDim.x) {
const auto xr = static_cast<const Float*>(pointer::offset<Float>(input, row * input_stride));
const auto yr = static_cast<Float*>(pointer::offset<Float>(output, row * output_stride));
float xi_cache[kElemsPerThread];
float lsq = 0.f;
#pragma unroll
for (int k = 0; k < kElemsPerThread; ++k) {
const int i = threadIdx.x + k * kWarpThreads;
xi_cache[k] = static_cast<float>(xr[i]);
lsq += xi_cache[k] * xi_cache[k];
}
lsq = warp::reduce_sum(lsq);
const float rstd = math::rsqrt(lsq / kDim + eps);
// HF semantics — round (x*rstd) to dtype, THEN multiply by weight.
#pragma unroll
for (int k = 0; k < kElemsPerThread; ++k) {
const int i = threadIdx.x + k * kWarpThreads;
const Float xn = cast<Float>(xi_cache[k] * rstd);
yr[i] = cast<Float>(static_cast<float>(xn) * static_cast<float>(wr[i]));
}
}
PDLTriggerSecondary<kUsePDL>();
}
// ---------------------------------------------------------------------------
// Kernel: 512-thread scalar-strided RMSNorm with HF semantics + register cache.
//
// Pass 1: each thread loads its strided elements, caches them in registers,
// and accumulates the fp32 sum-of-squares. Warp + block reduction
// yields `rstd = rsqrt(mean(x^2) + eps)`.
// Pass 2: reuse cached fp32 values — no second global read of `x`. Per-elem:
// xn = cast_to_dtype(x_fp32 * rstd) <- HF's cast-before-mul
// y = cast_to_dtype(float(xn) * float(w))
// ---------------------------------------------------------------------------
template <int64_t kDim, bool kUsePDL, typename Float>
__global__ __launch_bounds__(512) void rmsnorm_hf_scalar_kernel(const RMSNormHFParams __grid_constant__ params) {
using namespace device;
constexpr int kNumThreads = 512;
constexpr int kNumWarps = kNumThreads / kWarpThreads;
// For kDim=4096: kElemsPerThread = 8 (32 bytes of fp32 cache per thread).
constexpr int kElemsPerThread = (kDim + kNumThreads - 1) / kNumThreads;
const auto& [input, weight_ptr, output, input_stride, output_stride, num_tokens, eps] = params;
const auto xr = static_cast<const Float*>(pointer::offset<Float>(input, blockIdx.x * input_stride));
const auto yr = static_cast<Float*>(pointer::offset<Float>(output, blockIdx.x * output_stride));
const auto wr = static_cast<const Float*>(weight_ptr);
PDLWaitPrimary<kUsePDL>();
// Pass 1: load, square, accumulate; cache fp32 values in registers.
float xi_cache[kElemsPerThread];
float lsq = 0.f;
#pragma unroll
for (int k = 0; k < kElemsPerThread; ++k) {
const int i = threadIdx.x + k * kNumThreads;
xi_cache[k] = static_cast<float>(xr[i]);
lsq += xi_cache[k] * xi_cache[k];
}
// Warp reduce.
lsq = warp::reduce_sum(lsq);
// Block reduce via shared memory (32 warps * 1 fp32 each).
__shared__ float smem[32];
const int warp_id = threadIdx.x / kWarpThreads;
const int lane_id = threadIdx.x & (kWarpThreads - 1);
if (lane_id == 0) smem[warp_id] = lsq;
__syncthreads();
__shared__ float rstd_s;
if (threadIdx.x < kWarpThreads) {
float v = (threadIdx.x < kNumWarps) ? smem[threadIdx.x] : 0.f;
v = warp::reduce_sum(v);
if (threadIdx.x == 0) rstd_s = math::rsqrt(v / kDim + eps);
}
__syncthreads();
const float rstd = rstd_s;
// Pass 2: HF semantics — round (x*rstd) to dtype, THEN multiply by weight.
#pragma unroll
for (int k = 0; k < kElemsPerThread; ++k) {
const int i = threadIdx.x + k * kNumThreads;
const Float xn = cast<Float>(xi_cache[k] * rstd);
yr[i] = cast<Float>(static_cast<float>(xn) * static_cast<float>(wr[i]));
}
PDLTriggerSecondary<kUsePDL>();
}
// ---------------------------------------------------------------------------
// Warp launcher: occupancy-sized grid, 32 threads/block, one warp per row.
// Targets small hidden sizes (q/k RMSNorms). kDim must be a multiple of 32
// in [32, 512).
// ---------------------------------------------------------------------------
template <int64_t kDim, bool kUsePDL, typename DType>
struct HFRMSNormWarpKernel {
static_assert(sizeof(DType) == 2, "rmsnorm_hf: DType must be fp16_t or bf16_t");
static_assert(
kDim >= 32 && kDim < 512 && kDim % 32 == 0, "rmsnorm_hf_warp: kDim must be a multiple of 32, in [32, 512)");
static constexpr auto kernel = rmsnorm_hf_warp_kernel<kDim, kUsePDL, DType>;
static constexpr uint32_t kBlockSize = device::kWarpThreads;
static void
run(const tvm::ffi::TensorView input,
const tvm::ffi::TensorView weight,
const tvm::ffi::TensorView output,
float eps) {
using namespace host;
auto N = SymbolicSize{"num_tokens"};
auto D = SymbolicSize{"hidden_size"};
auto SI = SymbolicSize{"input_stride"};
auto SO = SymbolicSize{"output_stride"};
auto device_ = SymbolicDevice{};
D.set_value(kDim);
device_.set_options<kDLCUDA>();
TensorMatcher({N, D}).with_strides({SI, 1}).with_dtype<DType>().with_device(device_).verify(input);
TensorMatcher({D}).with_dtype<DType>().with_device(device_).verify(weight);
TensorMatcher({N, D}).with_strides({SO, 1}).with_dtype<DType>().with_device(device_).verify(output);
const auto num_tokens = static_cast<uint32_t>(N.unwrap());
RuntimeCheck(num_tokens > 0, "rmsnorm_hf: num_tokens must be > 0");
const auto params = RMSNormHFParams{
.input = input.data_ptr(),
.weight = weight.data_ptr(),
.output = output.data_ptr(),
.input_stride = SI.unwrap(),
.output_stride = SO.unwrap(),
.num_tokens = num_tokens,
.eps = eps,
};
static const uint32_t max_occupancy = runtime::get_blocks_per_sm(kernel, kBlockSize);
static const uint32_t kNumSM = runtime::get_sm_count(device_.unwrap().device_id);
const auto num_blocks = std::min<uint32_t>(num_tokens, max_occupancy * kNumSM);
LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) //
.enable_pdl(kUsePDL)(kernel, params);
}
};
// ---------------------------------------------------------------------------
// CTA launcher: validates tensors, launches one block per row.
// ---------------------------------------------------------------------------
template <int64_t kDim, bool kUsePDL, typename DType>
struct HFRMSNormKernel {
static_assert(sizeof(DType) == 2, "rmsnorm_hf: DType must be fp16_t or bf16_t");
static_assert(kDim >= 512 && kDim % 512 == 0, "rmsnorm_hf: kDim must be a multiple of 512");
static constexpr auto kernel = rmsnorm_hf_scalar_kernel<kDim, kUsePDL, DType>;
static constexpr uint32_t kBlockSize = 512;
static void
run(const tvm::ffi::TensorView input,
const tvm::ffi::TensorView weight,
const tvm::ffi::TensorView output,
float eps) {
using namespace host;
auto N = SymbolicSize{"num_tokens"};
auto D = SymbolicSize{"hidden_size"};
auto SI = SymbolicSize{"input_stride"};
auto SO = SymbolicSize{"output_stride"};
auto device_ = SymbolicDevice{};
D.set_value(kDim);
device_.set_options<kDLCUDA>();
TensorMatcher({N, D}) // input
.with_strides({SI, 1})
.with_dtype<DType>()
.with_device(device_)
.verify(input);
TensorMatcher({D}) // weight
.with_dtype<DType>()
.with_device(device_)
.verify(weight);
TensorMatcher({N, D}) // output
.with_strides({SO, 1})
.with_dtype<DType>()
.with_device(device_)
.verify(output);
const auto num_tokens = static_cast<uint32_t>(N.unwrap());
RuntimeCheck(num_tokens > 0, "rmsnorm_hf: num_tokens must be > 0");
const auto params = RMSNormHFParams{
.input = input.data_ptr(),
.weight = weight.data_ptr(),
.output = output.data_ptr(),
.input_stride = SI.unwrap(),
.output_stride = SO.unwrap(),
.num_tokens = num_tokens,
.eps = eps,
};
LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) //
.enable_pdl(kUsePDL)(kernel, params);
}
};
} // namespace
+79
View File
@@ -0,0 +1,79 @@
"""RMSNorm with HF LlamaRMSNorm semantics (cast to dtype before weight multiply)."""
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,
)
if TYPE_CHECKING:
from tvm_ffi.module import Module
_CTA_BLOCK_SIZE = 512
_WARP_SIZE = 32
def is_supported_rmsnorm_hf_hidden_size(hidden_size: int) -> bool:
"""Return True iff the JIT rmsnorm_hf kernel supports this hidden size.
Two launch configs cover the practical range:
- Warp kernel: ``[32, 512)`` in multiples of 32 (q/k RMSNorm head dims).
- CTA kernel: ``>= 512`` in multiples of 512 (token RMSNorms).
"""
if _WARP_SIZE <= hidden_size < _CTA_BLOCK_SIZE and hidden_size % _WARP_SIZE == 0:
return True
return hidden_size >= _CTA_BLOCK_SIZE and hidden_size % _CTA_BLOCK_SIZE == 0
@cache_once
def _jit_rmsnorm_hf_module(hidden_size: int, dtype: torch.dtype) -> Module:
args = make_cpp_args(hidden_size, is_arch_support_pdl(), dtype)
kernel_cls = (
"HFRMSNormWarpKernel" if hidden_size < _CTA_BLOCK_SIZE else "HFRMSNormKernel"
)
return load_jit(
"rmsnorm_hf",
*args,
cuda_files=["elementwise/rmsnorm_hf.cuh"],
cuda_wrappers=[("rmsnorm_hf", f"{kernel_cls}<{args}>::run")],
)
def rmsnorm_hf(
input: torch.Tensor,
weight: torch.Tensor,
eps: float = 1e-6,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""RMSNorm: ``out = weight * cast_dtype(rsqrt(mean(x^2) + eps) * x)``.
``input`` must be 2D ``(num_tokens, hidden_size)``; callers with
higher-rank tensors should reshape first. ``hidden_size`` must satisfy
:func:`is_supported_rmsnorm_hf_hidden_size`. Empty inputs return an empty
output without launching the kernel.
"""
if input.dtype not in (torch.float16, torch.bfloat16):
raise RuntimeError(f"rmsnorm_hf: input must be fp16 or bf16, got {input.dtype}")
if input.dim() != 2:
raise RuntimeError(f"rmsnorm_hf: input must be 2D, got {input.dim()}D")
hidden_size = input.size(-1)
if not is_supported_rmsnorm_hf_hidden_size(hidden_size):
raise RuntimeError(
f"rmsnorm_hf: unsupported hidden_size={hidden_size} "
f"(must be a multiple of {_WARP_SIZE} in [{_WARP_SIZE}, {_CTA_BLOCK_SIZE}) "
f"or a multiple of {_CTA_BLOCK_SIZE})"
)
if out is None:
out = torch.empty_like(input)
if input.numel() == 0:
return out
module = _jit_rmsnorm_hf_module(hidden_size, input.dtype)
module.rmsnorm_hf(input, weight, out, eps)
return out
@@ -0,0 +1,136 @@
"""Tests for the JIT rmsnorm_hf kernel (HF LlamaRMSNorm semantics)."""
import itertools
import sys
import pytest
import torch
from sglang.jit_kernel.rmsnorm_hf import (
is_supported_rmsnorm_hf_hidden_size,
rmsnorm_hf,
)
from sglang.jit_kernel.utils import get_ci_test_range
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, suite="stage-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
EPS = 1e-5
DEVICE = "cuda"
DTYPES = [torch.float16, torch.bfloat16]
def hf_rmsnorm_reference(x: torch.Tensor, w: torch.Tensor, eps: float) -> torch.Tensor:
"""HF LlamaRMSNorm: normalize fp32, cast normalized x to dtype, then multiply weight."""
x_fp32 = x.to(torch.float32)
variance = x_fp32.pow(2).mean(-1, keepdim=True)
x_normed = x_fp32 * torch.rsqrt(variance + eps)
return w * x_normed.to(x.dtype)
def sgl_rmsnorm_reference(x: torch.Tensor, w: torch.Tensor, eps: float) -> torch.Tensor:
"""Old sgl_kernel.rmsnorm semantics — weight multiply in fp32, cast at the end."""
x_fp32 = x.to(torch.float32)
variance = x_fp32.pow(2).mean(-1, keepdim=True)
x_normed = x_fp32 * torch.rsqrt(variance + eps)
return (x_normed * w.to(torch.float32)).to(x.dtype)
BS_LIST = get_ci_test_range(
[1, 2, 4, 7, 16, 64, 128, 512, 1024, 4096],
[1, 16, 1024],
)
HIDDEN_SIZE_LIST = get_ci_test_range(
# Warp-kernel shapes (q/k RMSNorm head dims) + CTA-kernel shapes.
[32, 64, 96, 128, 256, 512, 1024, 2048, 3072, 4096, 8192, 16384],
[128, 512, 4096, 16384],
)
@pytest.mark.parametrize(
"batch_size,hidden_size",
list(itertools.product(BS_LIST, HIDDEN_SIZE_LIST)),
)
@pytest.mark.parametrize("dtype", DTYPES)
def test_rmsnorm_hf_correctness(
batch_size: int, hidden_size: int, dtype: torch.dtype
) -> None:
torch.manual_seed(0)
x = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype)
w = torch.randn(hidden_size, device=DEVICE, dtype=dtype)
out = rmsnorm_hf(x, w, EPS)
ref = hf_rmsnorm_reference(x, w, EPS)
# Loose atol — the kernel's block-reduce order differs from PyTorch's
# `mean`, producing ~1 fp16 ULP of drift on some shapes.
# The SGL-semantics regression guard below is what catches the cast-order
# bug this PR fixes; it's reduction-order-invariant.
torch.testing.assert_close(out, ref, atol=1e-2, rtol=1e-2)
@pytest.mark.parametrize("dtype", DTYPES)
def test_rmsnorm_hf_out_param(dtype: torch.dtype) -> None:
torch.manual_seed(0)
x = torch.randn(8, 4096, device=DEVICE, dtype=dtype)
w = torch.randn(4096, device=DEVICE, dtype=dtype)
out = torch.empty_like(x)
result = rmsnorm_hf(x, w, EPS, out=out)
assert result.data_ptr() == out.data_ptr()
torch.testing.assert_close(
out, hf_rmsnorm_reference(x, w, EPS), atol=1e-2, rtol=1e-2
)
@pytest.mark.parametrize("dtype", DTYPES)
def test_rmsnorm_hf_matches_hf_not_sgl(dtype: torch.dtype) -> None:
"""Regression guard: kernel must follow HF (cast-before-mul), not the old
sgl_kernel.rmsnorm semantics (fp32-mul-then-cast). Reduction-order drift
prevents a bit-exact assert against HF, so instead assert the kernel is
strictly closer to HF than to the SGL reference."""
torch.manual_seed(0)
x = torch.randn(64, 4096, device=DEVICE, dtype=dtype)
w = torch.randn(4096, device=DEVICE, dtype=dtype)
out = rmsnorm_hf(x, w, EPS).float()
hf_ref = hf_rmsnorm_reference(x, w, EPS).float()
sgl_ref = sgl_rmsnorm_reference(x, w, EPS).float()
assert (sgl_ref - hf_ref).abs().max() > 0, "inputs don't exercise the difference"
diff_hf = (out - hf_ref).abs().max().item()
diff_sgl = (out - sgl_ref).abs().max().item()
assert (
diff_hf < diff_sgl
), f"kernel closer to SGL than HF (hf={diff_hf}, sgl={diff_sgl})"
def test_rmsnorm_hf_empty_input() -> None:
"""Empty input must short-circuit: the C++ launcher rejects num_tokens=0."""
x = torch.empty(0, 4096, device=DEVICE, dtype=torch.float16)
w = torch.randn(4096, device=DEVICE, dtype=torch.float16)
out = rmsnorm_hf(x, w, EPS)
assert out.shape == x.shape and out.numel() == 0
@pytest.mark.parametrize(
("hidden_size", "expected"),
[
(16, False),
(32, True),
(64, True),
(96, True),
(128, True),
(256, True),
(288, True),
(384, True),
(500, False),
(512, True),
(3072, True),
(4096, True),
(8192, True),
(4097, False),
],
)
def test_is_supported_hidden_size(hidden_size: int, expected: bool) -> None:
assert is_supported_rmsnorm_hf_hidden_size(hidden_size) is expected
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
+39
View File
@@ -83,6 +83,27 @@ elif _is_hip:
# Fallback: vllm not available, will use forward_native
_has_vllm_rms_norm = False
if _is_cuda:
# HF-semantics RMSNorm kernel (JIT-compiled). Used when `cast_x_before_out_mul=True`
# (the transformers backend path) to produce outputs that are numerically identical
# to HuggingFace `LlamaRMSNorm`: the cast from fp32 to the activation dtype happens
# BEFORE the weight multiply, so the multiply is done in the narrow dtype.
_jit_rmsnorm_hf_available = False
try:
from sglang.jit_kernel.rmsnorm_hf import (
is_supported_rmsnorm_hf_hidden_size,
)
from sglang.jit_kernel.rmsnorm_hf import rmsnorm_hf as _jit_rmsnorm_hf
_jit_rmsnorm_hf_available = True
except ImportError:
def is_supported_rmsnorm_hf_hidden_size(d: int) -> bool:
return False
_jit_rmsnorm_hf = None
logger = logging.getLogger(__name__)
if _is_npu:
@@ -199,6 +220,7 @@ class RMSNorm(MultiPlatformOp):
if is_batch_invariant_mode_enabled():
if (
residual is not None
or self.cast_x_before_out_mul
or get_global_server_args().rl_on_policy_target == "fsdp"
):
return self.forward_native(x, residual, post_residual_addition)
@@ -207,6 +229,23 @@ class RMSNorm(MultiPlatformOp):
self.weight.data,
self.variance_epsilon,
)
if self.cast_x_before_out_mul and residual is None:
# Use HF-semantics kernel (cast to dtype before weight multiply).
if (
_jit_rmsnorm_hf_available
and x.dtype in (torch.float16, torch.bfloat16)
and self.weight.data.dtype == x.dtype
and is_supported_rmsnorm_hf_hidden_size(x.shape[-1])
):
out = _jit_rmsnorm_hf(
x.contiguous(), self.weight.data, self.variance_epsilon
)
else:
# Fallback: pure-Python HF semantics (already implemented in forward_native).
out = self.forward_native(x, None, None)
if needs_reshape:
out = out.reshape(original_shape)
return out
if residual is not None:
# TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition.
# but right now we can only have hidden_states+(residual+post_residual_addition).
+3
View File
@@ -270,6 +270,9 @@ def replace_rms_norm_class(rms_norm: nn.Module, hidden_size: int) -> nn.Module:
kwargs["weight_dtype"] = weight_meta.dtype
else:
kwargs["has_weight"] = False
kwargs["cast_x_before_out_mul"] = (
True # match HF fp16-weight-multiply semantics
)
base_cls = RMSNorm
norm = base_cls(**kwargs)