diff --git a/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh b/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh index 0049f5258..7750d1266 100644 --- a/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh +++ b/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh @@ -53,6 +53,96 @@ SGL_DEVICE CacheDType load_cache_value(const CacheDType* ptr, int64_t idx) { #endif } +template +SGL_DEVICE T rotary_mul_rn(T lhs, T rhs) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 + uint16_t lhs_bits; + uint16_t rhs_bits; + if constexpr (std::is_same_v) { + lhs_bits = __bfloat16_as_ushort(lhs); + rhs_bits = __bfloat16_as_ushort(rhs); + } else { + lhs_bits = __half_as_ushort(lhs); + rhs_bits = __half_as_ushort(rhs); + } + uint16_t out_bits; + if constexpr (std::is_same_v) { + asm volatile("mul.rn.bf16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } else { + asm volatile("mul.rn.f16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } + if constexpr (std::is_same_v) { + return __ushort_as_bfloat16(out_bits); + } else { + return __ushort_as_half(out_bits); + } +#else + return lhs * rhs; +#endif +} + +template +SGL_DEVICE T rotary_add(T x, T cos, T y, T sin) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 + // nvcc may contract the packed local expression on SM120 even though the + // reference RoPE kernel rounds both products to the activation dtype first. + const T lhs = rotary_mul_rn(x, cos); + const T rhs = rotary_mul_rn(y, sin); + uint16_t lhs_bits; + uint16_t rhs_bits; + if constexpr (std::is_same_v) { + lhs_bits = __bfloat16_as_ushort(lhs); + rhs_bits = __bfloat16_as_ushort(rhs); + } else { + lhs_bits = __half_as_ushort(lhs); + rhs_bits = __half_as_ushort(rhs); + } + uint16_t out_bits; + if constexpr (std::is_same_v) { + asm volatile("add.rn.bf16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } else { + asm volatile("add.rn.f16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } + if constexpr (std::is_same_v) { + return __ushort_as_bfloat16(out_bits); + } else { + return __ushort_as_half(out_bits); + } +#else + return x * cos + y * sin; +#endif +} + +template +SGL_DEVICE T rotary_sub(T x, T cos, T y, T sin) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 + const T lhs = rotary_mul_rn(x, cos); + const T rhs = rotary_mul_rn(y, sin); + uint16_t lhs_bits; + uint16_t rhs_bits; + if constexpr (std::is_same_v) { + lhs_bits = __bfloat16_as_ushort(lhs); + rhs_bits = __bfloat16_as_ushort(rhs); + } else { + lhs_bits = __half_as_ushort(lhs); + rhs_bits = __half_as_ushort(rhs); + } + uint16_t out_bits; + if constexpr (std::is_same_v) { + asm volatile("sub.rn.bf16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } else { + asm volatile("sub.rn.f16 %0, %1, %2;" : "=h"(out_bits) : "h"(lhs_bits), "h"(rhs_bits)); + } + if constexpr (std::is_same_v) { + return __ushort_as_bfloat16(out_bits); + } else { + return __ushort_as_half(out_bits); + } +#else + return x * cos - y * sin; +#endif +} + template < int64_t kHeadDim, int64_t kRopeDim, @@ -135,8 +225,8 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParams __grid_constant__ const auto half_idx = (lane_id % kHalfRotaryLanes) * kElemsPerThread + 2 * j + i; const auto cos = load_cache_value(cos_ptr, half_idx); const auto sin = load_cache_value(sin_ptr, half_idx); - values[i] = lane_id < kHalfRotaryLanes ? values[i] * cos - partner_values[i] * sin - : values[i] * cos + partner_values[i] * sin; + values[i] = lane_id < kHalfRotaryLanes ? rotary_sub(values[i], cos, partner_values[i], sin) + : rotary_add(values[i], cos, partner_values[i], sin); } } } @@ -150,8 +240,8 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParams __grid_constant__ const auto sin = load_cache_value(sin_ptr, half_idx); const auto x = values[0]; const auto y = values[1]; - values[0] = x * cos - y * sin; - values[1] = y * cos + x * sin; + values[0] = rotary_sub(x, cos, y, sin); + values[1] = rotary_add(y, cos, x, sin); } } }