From 02a7fabdcaabc28f20735e2de4b8bc2eff621786 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Luka=20Govedi=C4=8D?= Date: Mon, 9 Mar 2026 16:45:50 -0400 Subject: [PATCH] Add support for non-contiguous input for rms-quant (dynamic & block) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Luka Govedič --- ...fused_layernorm_dynamic_per_token_quant.cu | 49 ++++++++++++------- .../fused_kernels/layernorm_utils.cuh | 31 ++++++------ 2 files changed, 46 insertions(+), 34 deletions(-) diff --git a/csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu b/csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu index b9a9b5cc7e4..bff593af64b 100644 --- a/csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu +++ b/csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu @@ -15,13 +15,13 @@ __device__ void rms_norm_dynamic_per_token_quant_vec( scalar_t const* __restrict__ input, // [..., hidden_size] scalar_t const* __restrict__ weight, // [hidden_size] float const* scale_ub, float const var_epsilon, int32_t const hidden_size, - scalar_t* __restrict__ residual = nullptr) { + int32_t const input_stride, scalar_t* __restrict__ residual = nullptr) { float rms = 0.0f; float token_scale = 0.0f; // Compute rms vllm::vectorized::compute_rms( - &rms, input, hidden_size, var_epsilon, residual); + &rms, input, hidden_size, input_stride, var_epsilon, residual); // Compute scale vllm::vectorized::compute_dynamic_per_token_scales) { token_scale = 1.0f / token_scale; vllm::vectorized::norm_and_quant( - out, input, weight, rms, &token_scale, hidden_size, residual); + has_residual>(out, input, weight, rms, + &token_scale, hidden_size, + input_stride, residual); } else { // FP8 - Do not invert token_scale for exact match with FBGemm vllm::vectorized::norm_and_quant( - out, input, weight, rms, &token_scale, hidden_size, residual); + has_residual>(out, input, weight, rms, + &token_scale, hidden_size, + input_stride, residual); } } @@ -51,7 +53,7 @@ __global__ void rms_norm_dynamic_per_token_quant_kernel( scalar_t const* __restrict__ input, // [..., hidden_size] scalar_t const* __restrict__ weight, // [hidden_size] float const* scale_ub, float const var_epsilon, int32_t const hidden_size, - scalar_t* __restrict__ residual = nullptr) { + int32_t const input_stride, scalar_t* __restrict__ residual = nullptr) { // For vectorization, token_input and token_output pointers need to be // aligned at 8-byte and 4-byte addresses respectively. bool const can_vectorize = hidden_size % 4 == 0; @@ -60,15 +62,15 @@ __global__ void rms_norm_dynamic_per_token_quant_kernel( return rms_norm_dynamic_per_token_quant_vec( out, scales, input, weight, scale_ub, var_epsilon, hidden_size, - residual); + input_stride, residual); } float rms = 0.0f; float token_scale = 0.0f; // Compute RMS - vllm::compute_rms(&rms, input, hidden_size, - var_epsilon, residual); + vllm::compute_rms( + &rms, input, hidden_size, input_stride, var_epsilon, residual); // Compute Scale vllm::compute_dynamic_per_token_scales( &token_scale, scales, input, weight, rms, scale_ub, hidden_size, @@ -78,11 +80,13 @@ __global__ void rms_norm_dynamic_per_token_quant_kernel( if constexpr (std::is_same_v) { token_scale = 1.0f / token_scale; vllm::norm_and_quant( - out, input, weight, rms, &token_scale, hidden_size, residual); + out, input, weight, rms, &token_scale, hidden_size, input_stride, + residual); } else { // FP8 - Do not invert s_token_scale for exact match with FBGemm vllm::norm_and_quant( - out, input, weight, rms, &token_scale, hidden_size, residual); + out, input, weight, rms, &token_scale, hidden_size, input_stride, + residual); } } @@ -97,12 +101,13 @@ __global__ void rms_norm_per_block_quant_kernel( scalar_t const* __restrict__ input, // [..., hidden_size] scalar_t const* __restrict__ weight, // [hidden_size] float const* scale_ub, float const var_epsilon, int32_t const hidden_size, - scalar_t* __restrict__ residual = nullptr, int64_t outer_scale_stride = 1) { + int32_t const input_stride, scalar_t* __restrict__ residual = nullptr, + int64_t outer_scale_stride = 1) { float rms; // Compute RMS // Always able to vectorize due to constraints on hidden_size vllm::vectorized::compute_rms( - &rms, input, hidden_size, var_epsilon, residual); + &rms, input, hidden_size, input_stride, var_epsilon, residual); // Compute Scale // Always able to vectorize due to constraints on hidden_size and group_size @@ -120,7 +125,7 @@ __global__ void rms_norm_per_block_quant_kernel( vllm::vectorized::norm_and_quant< scalar_t, scalar_out_t, std::is_same_v, has_residual, is_scale_transposed, group_size>( - out, input, weight, rms, scales, hidden_size, residual, + out, input, weight, rms, scales, hidden_size, input_stride, residual, outer_scale_stride); } @@ -137,6 +142,7 @@ void rms_norm_dynamic_per_token_quant_dispatch( std::optional const& scale_ub, std::optional& residual) { int32_t hidden_size = input.size(-1); + int32_t input_stride = input.view({-1, hidden_size}).stride(0); auto num_tokens = input.numel() / hidden_size; dim3 grid(num_tokens); @@ -153,7 +159,7 @@ void rms_norm_dynamic_per_token_quant_dispatch( out.data_ptr(), scales.data_ptr(), input.data_ptr(), weight.data_ptr(), scale_ub.has_value() ? scale_ub->data_ptr() : nullptr, - var_epsilon, hidden_size, + var_epsilon, hidden_size, input_stride, has_residual ? residual->data_ptr() : nullptr); }); }); @@ -170,7 +176,9 @@ void rms_norm_dynamic_per_token_quant( ? c10::ScalarType::Float8_e4m3fn : c10::ScalarType::Float8_e4m3fnuz; TORCH_CHECK(out.dtype() == kFp8Type || out.dtype() == torch::kInt8); - TORCH_CHECK(out.is_contiguous() && input.is_contiguous()); + TORCH_CHECK(out.is_contiguous()); + TORCH_CHECK(input.stride(-1) == 1, + "Input must be contiguous in the last dimension"); if (scale_ub.has_value()) { TORCH_CHECK(out.dtype() == kFp8Type); @@ -200,6 +208,7 @@ void rms_norm_per_block_quant_dispatch( std::optional const& scale_ub, std::optional& residual, bool is_scale_transposed) { int32_t hidden_size = input.size(-1); + int32_t input_stride = input.view({-1, hidden_size}).stride(0); auto num_tokens = input.numel() / hidden_size; dim3 grid(num_tokens); @@ -225,7 +234,7 @@ void rms_norm_per_block_quant_dispatch( weight.data_ptr(), scale_ub.has_value() ? scale_ub->data_ptr() : nullptr, - var_epsilon, hidden_size, + var_epsilon, hidden_size, input_stride, has_residual ? residual->data_ptr() : nullptr, scales.stride(1)); @@ -246,7 +255,9 @@ void rms_norm_per_block_quant(torch::Tensor& out, torch::Tensor const& input, ? c10::ScalarType::Float8_e4m3fn : c10::ScalarType::Float8_e4m3fnuz; TORCH_CHECK(out.dtype() == kFp8Type || out.dtype() == torch::kInt8); - TORCH_CHECK(out.is_contiguous() && input.is_contiguous()); + TORCH_CHECK(out.is_contiguous()); + TORCH_CHECK(input.stride(-1) == 1, + "Input must be contiguous in the last dimension"); if (scale_ub.has_value()) { TORCH_CHECK(out.dtype() == kFp8Type); diff --git a/csrc/quantization/fused_kernels/layernorm_utils.cuh b/csrc/quantization/fused_kernels/layernorm_utils.cuh index edf4024f0d4..0397c13d340 100644 --- a/csrc/quantization/fused_kernels/layernorm_utils.cuh +++ b/csrc/quantization/fused_kernels/layernorm_utils.cuh @@ -16,9 +16,10 @@ namespace vllm { // has_residual must be true, if residual is not a nullptr template __device__ void compute_rms(float* rms, scalar_t const* __restrict__ input, - int32_t const hidden_size, float const epsilon, + int32_t const hidden_size, + int32_t const input_stride, float const epsilon, scalar_t const* __restrict__ residual = nullptr) { - int64_t const token_offset = blockIdx.x * static_cast(hidden_size); + int64_t const token_offset = blockIdx.x * static_cast(input_stride); // sum of squares float ss = 0.0f; @@ -185,9 +186,10 @@ template (hidden_size); + int32_t const hidden_size, int32_t const input_stride, + scalar_t* __restrict__ residual = nullptr, int32_t const group_size = 0, + int64_t outer_scale_stride = 1) { + int64_t const token_offset = blockIdx.x * static_cast(input_stride); for (auto i = threadIdx.x; i < hidden_size; i += blockDim.x) { float x = static_cast(input[token_offset + i]); @@ -224,9 +226,10 @@ namespace vectorized { // hidden_size must be a multiple of 4 template __device__ void compute_rms(float* rms, scalar_t const* __restrict__ input, - int32_t const hidden_size, float const epsilon, + int32_t const hidden_size, + int32_t const input_stride, float const epsilon, scalar_t const* __restrict__ residual = nullptr) { - int64_t const token_offset = blockIdx.x * static_cast(hidden_size); + int64_t const token_offset = blockIdx.x * static_cast(input_stride); // Vectorized input/output to better utilize memory bandwidth. vec4_t const* vec_input = @@ -462,14 +465,12 @@ __device__ void compute_dynamic_per_token_scales( template -__device__ void norm_and_quant(scalar_out_t* __restrict__ output, - scalar_t const* __restrict__ input, - scalar_t const* __restrict__ weight, - float const rms, float* const scale, - int32_t const hidden_size, - scalar_t* __restrict__ residual = nullptr, - int64_t outer_scale_stride = 1) { - int64_t const token_offset = blockIdx.x * static_cast(hidden_size); +__device__ void norm_and_quant( + scalar_out_t* __restrict__ output, scalar_t const* __restrict__ input, + scalar_t const* __restrict__ weight, float const rms, float* const scale, + int32_t const hidden_size, int32_t const input_stride, + scalar_t* __restrict__ residual = nullptr, int64_t outer_scale_stride = 1) { + int64_t const token_offset = blockIdx.x * static_cast(input_stride); // Vectorized input/output/weight/residual to better utilize memory bandwidth. vec4_t const* vec_input =