From 02a1f23711c5bdbff81eb8a610dde39e1141d036 Mon Sep 17 00:00:00 2001 From: Giancarlo Delfin <32987265+TheEpicDolphin@users.noreply.github.com> Date: Thu, 25 Jun 2026 21:32:07 -0500 Subject: [PATCH] [DFlash] Fuse precompute kv per-layer rmsnorms (#46761) Signed-off-by: Giancarlo Delfin Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- csrc/libtorch_stable/layernorm_kernels.cu | 33 ++++++--- .../core/test_batched_weight_rms_norm.py | 70 +++++++++++++++++++ vllm/model_executor/models/qwen3_dflash.py | 23 +++--- 3 files changed, 108 insertions(+), 18 deletions(-) create mode 100644 tests/kernels/core/test_batched_weight_rms_norm.py diff --git a/csrc/libtorch_stable/layernorm_kernels.cu b/csrc/libtorch_stable/layernorm_kernels.cu index f29734fc265..0c59a09b1b9 100644 --- a/csrc/libtorch_stable/layernorm_kernels.cu +++ b/csrc/libtorch_stable/layernorm_kernels.cu @@ -20,20 +20,27 @@ __global__ void rms_norm_kernel( const int64_t input_stride_d4, // input.stride(-4) const int64_t input_shape_d2, // input.size(-2) const int64_t input_shape_d3, // input.size(-3) - const scalar_t* __restrict__ weight, // [hidden_size], null if !HasWeight + const scalar_t* __restrict__ weight, // [hidden_size] or + // [num_groups, hidden_size]; + // null if !HasWeight + const int64_t weight_stride, // 0 or weight.stride(0) const float epsilon, const int num_tokens, const int hidden_size) { __shared__ float s_variance; float variance = 0.0f; const scalar_t* input_row; + const scalar_t* weight_row; + int64_t weight_row_off = 0; if constexpr (NUM_DIMS == 2) { // 2D for layernorm normal case [batch_size, hidden] input_row = input + blockIdx.x * input_stride_d2; + weight_row = weight + blockIdx.x * weight_stride; } else if constexpr (NUM_DIMS == 3) { // 3D for q/k norm [batch_size, num_heads, head_size] int batch_idx = blockIdx.x / input_shape_d2; int head_idx = blockIdx.x % input_shape_d2; input_row = input + batch_idx * input_stride_d3 + head_idx * input_stride_d2; + weight_row = weight + batch_idx * weight_stride; } else if constexpr (NUM_DIMS == 4) { // 4D for transformers model_impl qk norm [batch, seq, head, head_dim] int batch_idx = blockIdx.x / (input_shape_d3 * input_shape_d2); @@ -42,6 +49,7 @@ __global__ void rms_norm_kernel( int head_idx = remaining % input_shape_d2; input_row = input + batch_idx * input_stride_d4 + seq_idx * input_stride_d3 + head_idx * input_stride_d2; + weight_row = weight + batch_idx * weight_stride; } auto vec_op = [&variance](const vec_n_t& vec) { @@ -69,7 +77,7 @@ __global__ void rms_norm_kernel( scalar_t* out_row = out + blockIdx.x * hidden_size; auto* v_in = reinterpret_cast*>(input_row); - auto* v_w = reinterpret_cast*>(weight); + auto* v_w = reinterpret_cast*>(weight_row); auto* v_out = reinterpret_cast*>(out_row); for (int i = threadIdx.x; i < hidden_size / VEC_SIZE; i += blockDim.x) { vec_n_t dst; @@ -211,15 +219,24 @@ fused_add_rms_norm_kernel( void rms_norm(torch::stable::Tensor& out, // [..., hidden_size] torch::stable::Tensor& input, // [..., hidden_size] - std::optional weight, // [hidden_size] - double epsilon) { + std::optional weight, double epsilon) { STD_TORCH_CHECK(out.is_contiguous()); if (input.stride(-1) != 1) { input = torch::stable::contiguous(input); } STD_TORCH_CHECK(input.stride(-1) == 1); + int64_t weight_stride = 0; if (weight.has_value()) { STD_TORCH_CHECK(weight->is_contiguous()); + if (weight->dim() == 1) { + STD_TORCH_CHECK(weight->size(0) == input.size(-1)); + } else if (weight->dim() == 2) { + STD_TORCH_CHECK(weight->size(0) == input.size(0)); + STD_TORCH_CHECK(weight->size(-1) == input.size(-1)); + weight_stride = weight->stride(0); + } else { + STD_TORCH_CHECK(false, "rms_norm weight must be 1D or 2D"); + } } int hidden_size = input.size(-1); @@ -256,16 +273,16 @@ void rms_norm(torch::stable::Tensor& out, // [..., hidden_size] out.mutable_data_ptr(), input.const_data_ptr(), input_stride_d2, input_stride_d3, input_stride_d4, input_shape_d2, - input_shape_d3, weight_ptr, epsilon, num_tokens, - hidden_size); + input_shape_d3, weight_ptr, weight_stride, epsilon, + num_tokens, hidden_size); } else { vllm::rms_norm_kernel <<>>( out.mutable_data_ptr(), input.const_data_ptr(), input_stride_d2, input_stride_d3, input_stride_d4, input_shape_d2, - input_shape_d3, weight_ptr, epsilon, num_tokens, - hidden_size); + input_shape_d3, weight_ptr, /*weight_stride=*/0, epsilon, + num_tokens, hidden_size); } }); }); diff --git a/tests/kernels/core/test_batched_weight_rms_norm.py b/tests/kernels/core/test_batched_weight_rms_norm.py new file mode 100644 index 00000000000..42711fdf09e --- /dev/null +++ b/tests/kernels/core/test_batched_weight_rms_norm.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Tests for the batched-weight RMS norm kernel (vllm._custom_ops.rms_norm). + +``rms_norm`` can use the outermost input batch index to select the corresponding +weight row. The result must match that of looping ``rms_norm`` over that dimension. +""" + +import pytest +import torch + +from vllm import _custom_ops as ops +from vllm.platforms import current_platform +from vllm.utils.torch_utils import set_random_seed + +pytestmark = pytest.mark.skipif( + not current_platform.is_cuda_alike(), + reason="rms_norm requires a CUDA/ROCm device", +) + + +@pytest.mark.parametrize( + "shape", + [ + (28, 17, 128), # 3D: [num_rows, tokens, hidden] + (1, 5, 2, 128), # 4D: single row (edge case) + (28, 13, 8, 128), # 4D: [L, num_ctx, nkv, hd] (DFlash K-norm) + (6, 3, 4, 769), # 4D: non-power-of-two hidden size + ], +) +@pytest.mark.parametrize("dtype", [torch.half, torch.bfloat16, torch.float]) +@pytest.mark.parametrize("seed", [42]) +@torch.inference_mode() +def test_rms_norm_matches_loop( + shape: tuple[int, ...], dtype: torch.dtype, seed: int +) -> None: + set_random_seed(seed) + torch.set_default_device("cuda") + + num_rows, hidden = shape[0], shape[-1] + eps = 1e-6 + + x = torch.randn(*shape, dtype=dtype) * 0.1 + # Distinct weight per row so that a wrong row index would be caught. + weight = torch.randn(num_rows, hidden, dtype=dtype) * 0.1 + 1.0 + + # Reference batched-weight rms norm. + out_ref = torch.empty_like(x) + for i in range(x.shape[0]): + ops.rms_norm(out_ref[i], x[i], weight[i], eps) + + out = torch.empty_like(x) + ops.rms_norm(out, x, weight, eps) + + # Expect bitwise-identical results. + torch.testing.assert_close(out, out_ref, atol=0, rtol=0) + + +@torch.inference_mode() +def test_rms_norm_validates_shapes() -> None: + torch.set_default_device("cuda") + + x = torch.randn(4, 8, 128, dtype=torch.float) + out = torch.empty_like(x) + # Expect num rows mismatch. + with pytest.raises(RuntimeError): + ops.rms_norm(out, x, torch.randn(3, 128), 1e-6) + # Expect hidden size mismatch. + with pytest.raises(RuntimeError): + ops.rms_norm(out, x, torch.randn(4, 64), 1e-6) diff --git a/vllm/model_executor/models/qwen3_dflash.py b/vllm/model_executor/models/qwen3_dflash.py index 36c0a357878..8746a15f115 100644 --- a/vllm/model_executor/models/qwen3_dflash.py +++ b/vllm/model_executor/models/qwen3_dflash.py @@ -309,8 +309,11 @@ class DFlashQwen3Model(nn.Module): else: self._fused_kv_bias = None - # K-norm weights: list of [head_dim] tensors, one per layer. - self._k_norm_weights = [a.k_norm.weight.data for a in layers_attn] + # K-norm weights stacked into one contiguous [num_layers, head_dim] + # tensor so the per-layer K-norm runs as a single grouped kernel. + self._k_norm_weights = torch.stack( + [a.k_norm.weight.data for a in layers_attn], dim=0 + ).contiguous() # RoPE parameters self._rope_head_size = attn0.rotary_emb.head_size @@ -392,15 +395,15 @@ class DFlashQwen3Model(nn.Module): all_k = all_kv[0] # [L, num_ctx, nkv, hd], contiguous all_v = all_kv[1] # [L, num_ctx, nkv, hd], contiguous - # --- Per-layer RMSNorm K (3D: [num_ctx, nkv, hd] per layer) --- + # --- Grouped RMSNorm K across all layers ([L, num_ctx, nkv, hd]) --- + # The weight is selected per layer by the outermost (layer) index. all_k_normed = torch.empty_like(all_k) - for i in range(L): - ops.rms_norm( - all_k_normed[i], - all_k[i], - self._k_norm_weights[i], - self._rms_norm_eps, - ) + ops.rms_norm( + all_k_normed, + all_k, + self._k_norm_weights, + self._rms_norm_eps, + ) # --- Fused RoPE across all layers --- # View as [L * num_ctx, kv] so RoPE sees one big batch (no copy).