[Diffusion][MiniMax H3] Extend exact QKNorm+RoPE rounding to SM103 (#34505)

This commit is contained in:
Xiaoyu Zhang
2026-08-12 16:25:24 +08:00
committed by GitHub
parent 45f7063335
commit 84ce7502cf
@@ -55,7 +55,7 @@ SGL_DEVICE CacheDType load_cache_value(const CacheDType* ptr, int64_t idx) {
template <typename T> template <typename T>
SGL_DEVICE T rotary_mul_rn(T lhs, T rhs) { SGL_DEVICE T rotary_mul_rn(T lhs, T rhs) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1030 || __CUDA_ARCH__ >= 1200)
uint16_t lhs_bits; uint16_t lhs_bits;
uint16_t rhs_bits; uint16_t rhs_bits;
if constexpr (std::is_same_v<T, bf16_t>) { if constexpr (std::is_same_v<T, bf16_t>) {
@@ -83,9 +83,10 @@ SGL_DEVICE T rotary_mul_rn(T lhs, T rhs) {
template <typename T> template <typename T>
SGL_DEVICE T rotary_add(T x, T cos, T y, T sin) { SGL_DEVICE T rotary_add(T x, T cos, T y, T sin) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1030 || __CUDA_ARCH__ >= 1200)
// nvcc may contract the packed local expression on SM120 even though the // nvcc may contract the packed local expression on Blackwell SM103/SM120
// reference RoPE kernel rounds both products to the activation dtype first. // even though the reference RoPE kernel rounds both products to the
// activation dtype first.
const T lhs = rotary_mul_rn(x, cos); const T lhs = rotary_mul_rn(x, cos);
const T rhs = rotary_mul_rn(y, sin); const T rhs = rotary_mul_rn(y, sin);
uint16_t lhs_bits; uint16_t lhs_bits;
@@ -115,7 +116,7 @@ SGL_DEVICE T rotary_add(T x, T cos, T y, T sin) {
template <typename T> template <typename T>
SGL_DEVICE T rotary_sub(T x, T cos, T y, T sin) { SGL_DEVICE T rotary_sub(T x, T cos, T y, T sin) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1200 #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1030 || __CUDA_ARCH__ >= 1200)
const T lhs = rotary_mul_rn(x, cos); const T lhs = rotary_mul_rn(x, cos);
const T rhs = rotary_mul_rn(y, sin); const T rhs = rotary_mul_rn(y, sin);
uint16_t lhs_bits; uint16_t lhs_bits;