[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:
@@ -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
|
||||||
@@ -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"]))
|
||||||
@@ -83,6 +83,27 @@ elif _is_hip:
|
|||||||
# Fallback: vllm not available, will use forward_native
|
# Fallback: vllm not available, will use forward_native
|
||||||
_has_vllm_rms_norm = False
|
_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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
@@ -199,6 +220,7 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if is_batch_invariant_mode_enabled():
|
if is_batch_invariant_mode_enabled():
|
||||||
if (
|
if (
|
||||||
residual is not None
|
residual is not None
|
||||||
|
or self.cast_x_before_out_mul
|
||||||
or get_global_server_args().rl_on_policy_target == "fsdp"
|
or get_global_server_args().rl_on_policy_target == "fsdp"
|
||||||
):
|
):
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
return self.forward_native(x, residual, post_residual_addition)
|
||||||
@@ -207,6 +229,23 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
self.weight.data,
|
self.weight.data,
|
||||||
self.variance_epsilon,
|
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:
|
if residual is not None:
|
||||||
# TODO: Ideally we want to have (hidden_states+residual)+post_residual_addition.
|
# 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).
|
# but right now we can only have hidden_states+(residual+post_residual_addition).
|
||||||
|
|||||||
@@ -270,6 +270,9 @@ def replace_rms_norm_class(rms_norm: nn.Module, hidden_size: int) -> nn.Module:
|
|||||||
kwargs["weight_dtype"] = weight_meta.dtype
|
kwargs["weight_dtype"] = weight_meta.dtype
|
||||||
else:
|
else:
|
||||||
kwargs["has_weight"] = False
|
kwargs["has_weight"] = False
|
||||||
|
kwargs["cast_x_before_out_mul"] = (
|
||||||
|
True # match HF fp16-weight-multiply semantics
|
||||||
|
)
|
||||||
base_cls = RMSNorm
|
base_cls = RMSNorm
|
||||||
norm = base_cls(**kwargs)
|
norm = base_cls(**kwargs)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user