diff --git a/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.cu b/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.cu index 9138604a9ba5..27848b86b75e 100644 --- a/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.cu +++ b/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.cu @@ -35,6 +35,46 @@ namespace kernels //////////////////////////////////////////////////////////////////////////////////////////////////// +namespace +{ + +template +struct FusedQKNormRopeTypeTraits; + +template <> +struct FusedQKNormRopeTypeTraits +{ + using PackedType = half2; + + __device__ static float toFloat(half value) + { + return __half2float(value); + } + + __device__ static float2 toFloat2(PackedType value) + { + return __half22float2(value); + } +}; + +template <> +struct FusedQKNormRopeTypeTraits<__nv_bfloat16> +{ + using PackedType = __nv_bfloat162; + + __device__ static float toFloat(__nv_bfloat16 value) + { + return __bfloat162float(value); + } + + __device__ static float2 toFloat2(PackedType value) + { + return __bfloat1622float2(value); + } +}; + +} // namespace + // Select the RoPE position id for a given rotary half-dim under interleaved mRoPE. // Mirrors MRotaryEmbedding.apply_interleaved_rope: section 1 (height) drives // dims {1,4,7,...} up to mrope_section1*3, section 2 (width) drives {2,5,8,...} @@ -78,7 +118,18 @@ __device__ __forceinline__ void storeHeadElements( } *reinterpret_cast(&out[offsetThread]) = vec; } - else // __nv_fp8_e4m3 + else if constexpr (std::is_same_v) + { + vec_T vec; +#pragma unroll + for (int i = 0; i < vecSize; i++) + { + half2 vals = __floats2half2_rn(elements[2 * i], elements[2 * i + 1]); + reinterpret_cast(*(reinterpret_cast(&vec) + i)) = vals; + } + *reinterpret_cast(&out[offsetThread]) = vec; + } + else { static_assert(numElemsPerThread % 2 == 0, "FP8 store expects an even element count per thread"); #pragma unroll @@ -146,26 +197,26 @@ __device__ __forceinline__ uint16_t quantizeMinimaxM3Fp4x4( return static_cast(fp32_vec_to_e2m1(quantValues)); } -// Perform per-head QK Norm and RoPE in a single kernel, reading a BF16 input and +// Perform per-head QK Norm and RoPE in a single kernel, reading a 16-bit input and // writing to a (possibly different-dtype) output buffer. // head_dim: the dimension of each head // interleave: interleave=!is_neox. -// OutT: output element type (__nv_bfloat16 or __nv_fp8_e4m3). -template +// OutT: output element type (__nv_bfloat16, half or __nv_fp8_e4m3). +template __global__ void fusedQKNormRopeKernel( - __nv_bfloat16 const* qkv_in, // Combined QKV input [num_tokens, (num_heads_q+num_heads_k+num_heads_v)*head_dim] - OutT* qkv_out, // Output buffer, same layout as qkv_in - int const num_heads_q, // Number of query heads - int const num_heads_k, // Number of key heads - int const num_heads_v, // Number of value heads - bool const process_v, // Whether to copy-cast V heads into qkv_out - int const rotary_dim, // Dimension for RoPE - float const eps, // Epsilon for RMS normalization - __nv_bfloat16 const* q_weight, // RMSNorm weights for query - __nv_bfloat16 const* k_weight, // RMSNorm weights for key - float const base, // Base for RoPE computation - int const* position_ids, // Position IDs for RoPE - int const num_tokens, // Number of tokens + InT const* qkv_in, // Combined QKV input [num_tokens, (num_heads_q+num_heads_k+num_heads_v)*head_dim] + OutT* qkv_out, // Output buffer, same layout as qkv_in + int const num_heads_q, // Number of query heads + int const num_heads_k, // Number of key heads + int const num_heads_v, // Number of value heads + bool const process_v, // Whether to copy-cast V heads into qkv_out + int const rotary_dim, // Dimension for RoPE + float const eps, // Epsilon for RMS normalization + InT const* q_weight, // RMSNorm weights for query + InT const* k_weight, // RMSNorm weights for key + float const base, // Base for RoPE computation + int const* position_ids, // Position IDs for RoPE + int const num_tokens, // Number of tokens // parameters for yarn float factor, // factor in rope_scaling in config.json. When it is not 1.0, it means the model is using yarn. float low, // threshold for high frequency @@ -226,10 +277,12 @@ __global__ void fusedQKNormRopeKernel( "elements)"); constexpr int numElemsPerThread = head_dim / 32; float elements[numElemsPerThread]; - constexpr int elemSizeBytes = numElemsPerThread * sizeof(__nv_bfloat16); + constexpr int elemSizeBytes = numElemsPerThread * sizeof(InT); static_assert(elemSizeBytes % 4 == 0, "numSizeBytes must be a multiple of 4"); constexpr int vecSize = elemSizeBytes / 4; // Use packed_as to perform loading/saving. using vec_T = typename tensorrt_llm::common::packed_as::type; + using PackedType = typename FusedQKNormRopeTypeTraits::PackedType; + static_assert(sizeof(PackedType) == sizeof(uint)); int const offsetWarp = tokenIdx * num_heads * head_dim + segStart + headIdx * head_dim; int offsetThread = offsetWarp + laneId * numElemsPerThread; @@ -243,7 +296,8 @@ __global__ void fusedQKNormRopeKernel( #pragma unroll for (int i = 0; i < vecSize; i++) { - float2 vals = __bfloat1622float2(*reinterpret_cast<__nv_bfloat162*>(reinterpret_cast(&vec) + i)); + float2 vals = FusedQKNormRopeTypeTraits::toFloat2( + *reinterpret_cast(reinterpret_cast(&vec) + i)); sumOfSquares += vals.x * vals.x; sumOfSquares += vals.y * vals.y; @@ -271,7 +325,8 @@ __global__ void fusedQKNormRopeKernel( for (int i = 0; i < numElemsPerThread; i++) { int dim = laneId * numElemsPerThread + i; - float weight = isQ ? __bfloat162float(q_weight[dim]) : __bfloat162float(k_weight[dim]); + float weight = isQ ? FusedQKNormRopeTypeTraits::toFloat(q_weight[dim]) + : FusedQKNormRopeTypeTraits::toFloat(k_weight[dim]); // Gemma RMSNorm scales by (1 + weight); standard RMSNorm scales by weight. elements[i] *= rms_rcp * (use_gemma ? (1.0f + weight) : weight); } @@ -876,13 +931,12 @@ __global__ void minimaxM3Nvfp4QKVIndexerNormRopeKVInsertKernel(__nv_bfloat16 con __VA_ARGS__ \ } -template -static void launchFusedQKNormRopeImpl(__nv_bfloat16 const* qkv_in, OutT* qkv_out, bool const process_v, +template +static void launchFusedQKNormRopeImpl(InT const* qkv_in, OutT* qkv_out, bool const process_v, int const num_tokens, int const num_heads_q, int const num_heads_k, int const num_heads_v, int const head_dim, - int const rotary_dim, float const eps, __nv_bfloat16 const* q_weight, __nv_bfloat16 const* k_weight, - float const base, bool const interleave, int const* position_ids, float factor, float low, float high, - float attention_factor, cudaStream_t stream, bool is_qk_norm, bool use_gemma, bool use_mrope, int mrope_section1, - int mrope_section2) + int const rotary_dim, float const eps, InT const* q_weight, InT const* k_weight, float const base, + bool const interleave, int const* position_ids, float factor, float low, float high, float attention_factor, + cudaStream_t stream, bool is_qk_norm, bool use_gemma, bool use_mrope, int mrope_section1, int mrope_section2) { if (factor == 1.0f) { @@ -918,26 +972,26 @@ static void launchFusedQKNormRopeImpl(__nv_bfloat16 const* qkv_in, OutT* qkv_out { case 64: DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { - fusedQKNormRopeKernel<64, INTERLEAVE, OutT><<>>(qkv_in, qkv_out, num_heads_q, - num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, position_ids, - num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, mrope_section1, - mrope_section2); + fusedQKNormRopeKernel<<>>(qkv_in, qkv_out, + num_heads_q, num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, + position_ids, num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, + mrope_section1, mrope_section2); }); break; case 128: DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { - fusedQKNormRopeKernel<128, INTERLEAVE, OutT><<>>(qkv_in, qkv_out, num_heads_q, - num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, position_ids, - num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, mrope_section1, - mrope_section2); + fusedQKNormRopeKernel<<>>(qkv_in, qkv_out, + num_heads_q, num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, + position_ids, num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, + mrope_section1, mrope_section2); }); break; case 256: DISPATCH_INTERLEAVE(interleave, INTERLEAVE, { - fusedQKNormRopeKernel<256, INTERLEAVE, OutT><<>>(qkv_in, qkv_out, num_heads_q, - num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, position_ids, - num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, mrope_section1, - mrope_section2); + fusedQKNormRopeKernel<<>>(qkv_in, qkv_out, + num_heads_q, num_heads_k, num_heads_v, process_v, rotary_dim, eps, q_weight, k_weight, base, + position_ids, num_tokens, factor, low, high, attention_factor, is_qk_norm, use_gemma, use_mrope, + mrope_section1, mrope_section2); }); break; default: TLLM_THROW("Unsupported head dimension for fusedQKNormRope: %d", head_dim); @@ -946,15 +1000,26 @@ static void launchFusedQKNormRopeImpl(__nv_bfloat16 const* qkv_in, OutT* qkv_out void launchFusedQKNormRope(void* qkv, int const num_tokens, int const num_heads_q, int const num_heads_k, int const num_heads_v, int const head_dim, int const rotary_dim, float const eps, void const* q_weight, - void const* k_weight, float const base, bool const interleave, int const* position_ids, float factor, float low, - float high, float attention_factor, cudaStream_t stream, bool is_qk_norm, bool use_gemma, bool use_mrope, - int mrope_section1, int mrope_section2) + void const* k_weight, bool is_bfloat16, float const base, bool const interleave, int const* position_ids, + float factor, float low, float high, float attention_factor, cudaStream_t stream, bool is_qk_norm, bool use_gemma, + bool use_mrope, int mrope_section1, int mrope_section2) { - launchFusedQKNormRopeImpl<__nv_bfloat16>(static_cast<__nv_bfloat16 const*>(qkv), static_cast<__nv_bfloat16*>(qkv), - /*process_v=*/false, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim, rotary_dim, eps, - static_cast<__nv_bfloat16 const*>(q_weight), static_cast<__nv_bfloat16 const*>(k_weight), base, interleave, - position_ids, factor, low, high, attention_factor, stream, is_qk_norm, use_gemma, use_mrope, mrope_section1, - mrope_section2); + if (is_bfloat16) + { + launchFusedQKNormRopeImpl<__nv_bfloat16, __nv_bfloat16>(static_cast<__nv_bfloat16 const*>(qkv), + static_cast<__nv_bfloat16*>(qkv), /*process_v=*/false, num_tokens, num_heads_q, num_heads_k, num_heads_v, + head_dim, rotary_dim, eps, static_cast<__nv_bfloat16 const*>(q_weight), + static_cast<__nv_bfloat16 const*>(k_weight), base, interleave, position_ids, factor, low, high, + attention_factor, stream, is_qk_norm, use_gemma, use_mrope, mrope_section1, mrope_section2); + } + else + { + launchFusedQKNormRopeImpl(static_cast(qkv), static_cast(qkv), + /*process_v=*/false, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim, rotary_dim, eps, + static_cast(q_weight), static_cast(k_weight), base, interleave, position_ids, + factor, low, high, attention_factor, stream, is_qk_norm, use_gemma, use_mrope, mrope_section1, + mrope_section2); + } } void launchFusedQKNormRopeToFp8(void const* qkv_in, void* qkv_out, int const num_tokens, int const num_heads_q, @@ -964,7 +1029,7 @@ void launchFusedQKNormRopeToFp8(void const* qkv_in, void* qkv_out, int const num bool use_mrope, int mrope_section1, int mrope_section2) { // Out-of-place, so V has to be copy-cast rather than left untouched. - launchFusedQKNormRopeImpl<__nv_fp8_e4m3>(static_cast<__nv_bfloat16 const*>(qkv_in), + launchFusedQKNormRopeImpl<__nv_bfloat16, __nv_fp8_e4m3>(static_cast<__nv_bfloat16 const*>(qkv_in), static_cast<__nv_fp8_e4m3*>(qkv_out), /*process_v=*/true, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim, rotary_dim, eps, static_cast<__nv_bfloat16 const*>(q_weight), static_cast<__nv_bfloat16 const*>(k_weight), base, interleave, position_ids, factor, low, high, diff --git a/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.h b/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.h index 30f63edfd0c2..5928046ba617 100644 --- a/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.h +++ b/cpp/tensorrt_llm/kernels/fusedQKNormRopeKernel.h @@ -39,6 +39,7 @@ void launchFusedQKNormRope( float const eps, // Epsilon for RMS normalization void const* q_weight, // RMSNorm weights for query [head_dim] void const* k_weight, // RMSNorm weights for key [head_dim] + bool is_bfloat16, // Whether QKV and weights use bfloat16 (otherwise float16) float const base, // Base for RoPE computation bool const interleave, // Whether RoPE is applied in interleave mode (non-Neox style) int const* position_ids, // Position IDs for RoPE [num_tokens] diff --git a/cpp/tensorrt_llm/thop/fusedQKNormRopeOp.cpp b/cpp/tensorrt_llm/thop/fusedQKNormRopeOp.cpp index 4b1023003d7f..59f10326d605 100644 --- a/cpp/tensorrt_llm/thop/fusedQKNormRopeOp.cpp +++ b/cpp/tensorrt_llm/thop/fusedQKNormRopeOp.cpp @@ -47,10 +47,17 @@ int64_t validateFusedQKNormRopeInputs(torch::Tensor const& qkv, torch::Tensor co TORCH_CHECK(q_weight.size(0) == head_dim, "Query weights size must match head dimension"); TORCH_CHECK(k_weight.size(0) == head_dim, "Key weights size must match head dimension"); - CHECK_INPUT(qkv, torch::kBFloat16); + CHECK_TH_CUDA(qkv); + CHECK_CONTIGUOUS(qkv); + TORCH_CHECK(qkv.scalar_type() == torch::kFloat16 || qkv.scalar_type() == torch::kBFloat16, + "QKV tensor must have float16 or bfloat16 dtype"); CHECK_INPUT(position_ids, torch::kInt32); - CHECK_INPUT(q_weight, torch::kBFloat16); - CHECK_INPUT(k_weight, torch::kBFloat16); + CHECK_TH_CUDA(q_weight); + CHECK_CONTIGUOUS(q_weight); + CHECK_TYPE(q_weight, qkv.scalar_type()); + CHECK_TH_CUDA(k_weight); + CHECK_CONTIGUOUS(k_weight); + CHECK_TYPE(k_weight, qkv.scalar_type()); int64_t num_tokens = qkv.size(0); TORCH_CHECK(position_ids.size(-1) == num_tokens, "Number of tokens in position_ids must match QKV"); @@ -167,11 +174,11 @@ void fused_qk_norm_rope( auto stream = at::cuda::getCurrentCUDAStream(qkv.get_device()); - tensorrt_llm::kernels::launchFusedQKNormRope(reinterpret_cast<__nv_bfloat16*>(qkv.data_ptr()), - static_cast(num_tokens), static_cast(num_heads_q), static_cast(num_heads_k), - static_cast(num_heads_v), static_cast(head_dim), static_cast(rotary_dim), - static_cast(eps), reinterpret_cast<__nv_bfloat16*>(q_weight.data_ptr()), - reinterpret_cast<__nv_bfloat16*>(k_weight.data_ptr()), static_cast(base), + bool const isBfloat16 = qkv.scalar_type() == torch::kBFloat16; + tensorrt_llm::kernels::launchFusedQKNormRope(qkv.data_ptr(), static_cast(num_tokens), + static_cast(num_heads_q), static_cast(num_heads_k), static_cast(num_heads_v), + static_cast(head_dim), static_cast(rotary_dim), static_cast(eps), q_weight.data_ptr(), + k_weight.data_ptr(), isBfloat16, static_cast(base), !is_neox, // interleave reinterpret_cast(position_ids.data_ptr()), static_cast(factor), static_cast(low), static_cast(high), static_cast(attention_factor), stream, is_qk_norm, use_gemma, use_mrope, @@ -186,6 +193,7 @@ torch::Tensor fused_qk_norm_rope_to_fp8(torch::Tensor const& qkv, // [num_tokens torch::Tensor const& position_ids, double factor, double low, double high, double attention_factor, bool is_qk_norm, bool use_gemma, bool use_mrope, int64_t mrope_section1, int64_t mrope_section2) { + TORCH_CHECK(qkv.scalar_type() == torch::kBFloat16, "FP8 output requires a bfloat16 QKV input"); int64_t num_tokens = validateFusedQKNormRopeInputs( qkv, position_ids, q_weight, k_weight, num_heads_q, num_heads_k, num_heads_v, head_dim, use_mrope); diff --git a/tensorrt_llm/_torch/models/modeling_dflash.py b/tensorrt_llm/_torch/models/modeling_dflash.py index b9693465b6fe..c90ec94964a8 100644 --- a/tensorrt_llm/_torch/models/modeling_dflash.py +++ b/tensorrt_llm/_torch/models/modeling_dflash.py @@ -1506,7 +1506,8 @@ def dflash_forward( has_qk_norm = self._has_qk_norm is_bf16 = noise_embedding.dtype == torch.bfloat16 - use_fused_qk_norm_rope = self._use_fused_qk_norm_rope and is_bf16 + is_supported_dtype = noise_embedding.dtype in (torch.float16, torch.bfloat16) + use_fused_qk_norm_rope = self._use_fused_qk_norm_rope and is_supported_dtype use_fused_rope = ( _flashinfer_rope is not None and has_qk_norm and is_bf16 and not use_fused_qk_norm_rope ) diff --git a/tests/unittest/_torch/attention/kernels/parallel_hw_agnostic/test_fused_qk_norm_rope.py b/tests/unittest/_torch/attention/kernels/parallel_hw_agnostic/test_fused_qk_norm_rope.py index b8db4ae76479..1f96187e3207 100644 --- a/tests/unittest/_torch/attention/kernels/parallel_hw_agnostic/test_fused_qk_norm_rope.py +++ b/tests/unittest/_torch/attention/kernels/parallel_hw_agnostic/test_fused_qk_norm_rope.py @@ -213,7 +213,7 @@ def torch_ref_rms_norm_rope( num_tokens_list = [1, 3, 8, 32, 256] is_neox_list = [False, True] partial_rotary_factor_list = [1.0, 0.5] -dtypes = [torch.bfloat16] # TODO: support float16 +dtypes = [torch.bfloat16, torch.float16] @pytest.mark.parametrize("head_dim", head_dims) @@ -315,6 +315,8 @@ def test_fused_qk_norm_rope( rtol=5e-2, atol=1e-1, ) + v_offset = (num_heads_q + num_heads_k) * head_dim + torch.testing.assert_close(output[:, v_offset:], qkv_copy[:, v_offset:], rtol=0, atol=0) @torch.inference_mode()