From f9876d6a5bae19534d53eb90a1263544eea78f37 Mon Sep 17 00:00:00 2001 From: mingfeima Date: Wed, 15 Jul 2026 16:40:09 +0000 Subject: [PATCH 1/5] replace fp32 div with rcp14 and simplify code --- sgl-kernel/csrc/cpu/activation.cpp | 19 ++--- sgl-kernel/csrc/cpu/conv3d.cpp | 9 +-- sgl-kernel/csrc/cpu/decode.cpp | 8 +- sgl-kernel/csrc/cpu/flash_attn.h | 8 +- sgl-kernel/csrc/cpu/gemm.cpp | 25 ++---- sgl-kernel/csrc/cpu/gemm_fp8.cpp | 14 ++-- sgl-kernel/csrc/cpu/mamba/conv.cpp | 5 +- sgl-kernel/csrc/cpu/mamba/fla.cpp | 53 ++++-------- sgl-kernel/csrc/cpu/moe.cpp | 124 ++++------------------------ sgl-kernel/csrc/cpu/moe.h | 126 ++++++++++------------------- sgl-kernel/csrc/cpu/moe_int8.cpp | 30 +++---- sgl-kernel/csrc/cpu/qkv_proj.cpp | 12 +-- sgl-kernel/csrc/cpu/vec.h | 63 +++++++++++++-- 13 files changed, 181 insertions(+), 315 deletions(-) diff --git a/sgl-kernel/csrc/cpu/activation.cpp b/sgl-kernel/csrc/cpu/activation.cpp index da0951498e9b..8263e4ebe70e 100644 --- a/sgl-kernel/csrc/cpu/activation.cpp +++ b/sgl-kernel/csrc/cpu/activation.cpp @@ -25,13 +25,8 @@ void act_and_mul_kernel_impl( int64_t d; #pragma GCC unroll 4 for (d = 0; d <= dim - kVecSize; d += kVecSize) { - bVec x_bvec = bVec::loadu(input_ptr + d); - fVec x_fvec0, x_fvec1; - std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); - - bVec y_bvec = bVec::loadu(input_other_ptr + d); - fVec y_fvec0, y_fvec1; - std::tie(y_fvec0, y_fvec1) = at::vec::convert_to_float(y_bvec); + auto [x_fvec0, x_fvec1] = load_float_vec2(input_ptr + d); + auto [y_fvec0, y_fvec1] = load_float_vec2(input_other_ptr + d); x_fvec0 = vf(x_fvec0); x_fvec1 = vf(x_fvec1); @@ -39,8 +34,7 @@ void act_and_mul_kernel_impl( x_fvec0 = x_fvec0 * y_fvec0; x_fvec1 = x_fvec1 * y_fvec1; - x_bvec = convert_from_float_ext(x_fvec0, x_fvec1); - x_bvec.store(output_ptr + d); + convert_from_float_ext(x_fvec0, x_fvec1).store(output_ptr + d); } #pragma GCC unroll 4 for (; d < dim; ++d) { @@ -69,7 +63,6 @@ void fused_sigmoid_mul_kernel_impl( using fVec = at::vec::Vectorized; constexpr int64_t kVecSize = bVec::size(); - const fVec one = fVec(1.f); at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) { for (int64_t i = begin; i < end; ++i) { const scalar_t* __restrict__ i_ptr = input + i * dim; @@ -86,8 +79,8 @@ void fused_sigmoid_mul_kernel_impl( for (; d <= head_dim - kVecSize; d += kVecSize) { auto [x_fvec0, x_fvec1] = load_float_vec2(attn_ptr + d); auto [g_fvec0, g_fvec1] = load_float_vec2(gate_ptr + d); - x_fvec0 = x_fvec0 / (one + g_fvec0.neg().exp_u20()); - x_fvec1 = x_fvec1 / (one + g_fvec1.neg().exp_u20()); + x_fvec0 = x_fvec0 * fast_sigmoid(g_fvec0); + x_fvec1 = x_fvec1 * fast_sigmoid(g_fvec1); convert_from_float_ext(x_fvec0, x_fvec1).store(out_ptr + d); } #pragma GCC unroll 4 @@ -121,7 +114,7 @@ at::Tensor silu_and_mul_cpu(at::Tensor& input) { num_tokens, d, [](float x) { return x / (1.f + std::exp(-x)); }, - [](Vec x) { return x / (Vec(1.f) + x.neg().exp_u20()); }); + [](Vec x) { return fast_silu(x); }); }); return out; } diff --git a/sgl-kernel/csrc/cpu/conv3d.cpp b/sgl-kernel/csrc/cpu/conv3d.cpp index 4342e789c8cb..18b53e84121b 100644 --- a/sgl-kernel/csrc/cpu/conv3d.cpp +++ b/sgl-kernel/csrc/cpu/conv3d.cpp @@ -59,13 +59,12 @@ inline void copy_add_stub( constexpr int kVecSize = bVec::size(); for (int64_t d = 0; d < N; d += kVecSize) { - fVec bias0, bias1; - bVec bias_vec = bVec::loadu(bias + d); - std::tie(bias0, bias1) = at::vec::convert_to_float(bias_vec); + auto [bias0, bias1] = load_float_vec2(bias + d); for (int64_t m = 0; m < M; ++m) { - fVec data0 = fVec::loadu(Ctmp + m * N + d) + bias0; - fVec data1 = fVec::loadu(Ctmp + m * N + d + fVec::size()) + bias1; + auto [data0, data1] = load_float_vec2(Ctmp + m * N + d); + data0 = data0 + bias0; + data1 = data1 + bias1; bVec out_vec = convert_from_float_ext(data0, data1); out_vec.store(C + m * ldc + d); } diff --git a/sgl-kernel/csrc/cpu/decode.cpp b/sgl-kernel/csrc/cpu/decode.cpp index e8ca4554e37d..d2c338c48a00 100644 --- a/sgl-kernel/csrc/cpu/decode.cpp +++ b/sgl-kernel/csrc/cpu/decode.cpp @@ -150,8 +150,9 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, int64_t d = 0; #pragma GCC unroll 4 for (; d <= size - kVecSize; d += kVecSize) { - fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec; - fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec; + auto [a_fvec0, a_fvec1] = load_float_vec2(acc + d); + a_fvec0 = a_fvec0 * s_fvec; + a_fvec1 = a_fvec1 * s_fvec; bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); out_bvec.store(out + d); } @@ -186,8 +187,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu constexpr int col = i % COLS; // for COLS = 2, 4 use 512bit store if constexpr (col % 2 == 0) { - fVec a_fvec0 = fVec::loadu(input + col * 16); - fVec a_fvec1 = fVec::loadu(input + col * 16 + 16); + auto [a_fvec0, a_fvec1] = load_float_vec2(input + col * 16); bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); out_bvec.store(out + col * 16); } diff --git a/sgl-kernel/csrc/cpu/flash_attn.h b/sgl-kernel/csrc/cpu/flash_attn.h index ad90e989efda..95fbfcaf252e 100644 --- a/sgl-kernel/csrc/cpu/flash_attn.h +++ b/sgl-kernel/csrc/cpu/flash_attn.h @@ -29,8 +29,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu constexpr int col = i % COLS; // for COLS = 2, 4 use 512bit store if constexpr (col % 2 == 0) { - fVec a_fvec0 = fVec::loadu(input + col * 16); - fVec a_fvec1 = fVec::loadu(input + col * 16 + 16); + auto [a_fvec0, a_fvec1] = load_float_vec2(input + col * 16); bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); out_bvec.store(out + col * 16); } @@ -47,8 +46,9 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, int d = 0; #pragma GCC unroll 4 for (; d <= size - kVecSize; d += kVecSize) { - fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec; - fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec; + auto [a_fvec0, a_fvec1] = load_float_vec2(acc + d); + a_fvec0 = a_fvec0 * s_fvec; + a_fvec1 = a_fvec1 * s_fvec; bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); out_bvec.store(out + d); } diff --git a/sgl-kernel/csrc/cpu/gemm.cpp b/sgl-kernel/csrc/cpu/gemm.cpp index e66edacd49c2..904ef39f657e 100644 --- a/sgl-kernel/csrc/cpu/gemm.cpp +++ b/sgl-kernel/csrc/cpu/gemm.cpp @@ -111,8 +111,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d); - fVec data1 = fVec::loadu(input + d + fVec::size()); + auto [data0, data1] = load_float_vec2(input + d); bVec out_vec = convert_from_float_ext(data0, data1); out_vec.store(out + d); } @@ -130,9 +129,7 @@ inline void copy_stub(float* __restrict__ out, const scalar_t* __restrict__ inpu int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0, data1; - bVec b_vec = bVec::loadu(input + d); - std::tie(data0, data1) = at::vec::convert_to_float(b_vec); + auto [data0, data1] = load_float_vec2(input + d); data0.store(out + d); data1.store(out + d + fVec::size()); } @@ -151,9 +148,9 @@ inline void copy_add_stub( int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d) + fVec::loadu(bias + d); - fVec data1 = fVec::loadu(input + d + fVec::size()) + fVec::loadu(bias + d + fVec::size()); - bVec out_vec = convert_from_float_ext(data0, data1); + auto [data0, data1] = load_float_vec2(input + d); + auto [bias0, bias1] = load_float_vec2(bias + d); + bVec out_vec = convert_from_float_ext(data0 + bias0, data1 + bias1); out_vec.store(out + d); } for (; d < size; ++d) { @@ -171,7 +168,6 @@ inline void scalar_sigmoid_and_mul( using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; // scalar sigmoid - const fVec one = fVec(1.f); fVec X; if constexpr (has_bias) { assert(bias != nullptr); @@ -179,18 +175,13 @@ inline void scalar_sigmoid_and_mul( } else { X = fVec(input[0]); } - X = one / (one + X.neg().exp_u20()); + X = fast_sigmoid(X); // vec mul constexpr int kVecSize = bVec::size(); for (int d = 0; d < SIZE; d += kVecSize) { - bVec m_bvec = bVec::loadu(mul + d); - fVec m_fvec0, m_fvec1; - std::tie(m_fvec0, m_fvec1) = at::vec::convert_to_float(m_bvec); - m_fvec0 = m_fvec0 * X; - m_fvec1 = m_fvec1 * X; - - bVec out_vec = convert_from_float_ext(m_fvec0, m_fvec1); + auto [m_fvec0, m_fvec1] = load_float_vec2(mul + d); + bVec out_vec = convert_from_float_ext(m_fvec0 * X, m_fvec1 * X); out_vec.store(out + d); } } diff --git a/sgl-kernel/csrc/cpu/gemm_fp8.cpp b/sgl-kernel/csrc/cpu/gemm_fp8.cpp index c176c418d7f5..06b0b7f133e7 100644 --- a/sgl-kernel/csrc/cpu/gemm_fp8.cpp +++ b/sgl-kernel/csrc/cpu/gemm_fp8.cpp @@ -13,8 +13,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d); - fVec data1 = fVec::loadu(input + d + fVec::size()); + auto [data0, data1] = load_float_vec2(input + d); bVec out_vec = convert_from_float_ext(data0, data1); out_vec.store(out + d); } @@ -33,9 +32,9 @@ inline void copy_add_stub( int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d) + fVec::loadu(bias + d); - fVec data1 = fVec::loadu(input + d + fVec::size()) + fVec::loadu(bias + d + fVec::size()); - bVec out_vec = convert_from_float_ext(data0, data1); + auto [data0, data1] = load_float_vec2(input + d); + auto [bias0, bias1] = load_float_vec2(bias + d); + bVec out_vec = convert_from_float_ext(data0 + bias0, data1 + bias1); out_vec.store(out + d); } for (; d < size; ++d) { @@ -52,9 +51,8 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ int d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d) * vscale; - fVec data1 = fVec::loadu(input + d + fVec::size()) * vscale; - bVec out_vec = convert_from_float_ext(data0, data1); + auto [data0, data1] = load_float_vec2(input + d); + bVec out_vec = convert_from_float_ext(data0 * vscale, data1 * vscale); out_vec.store(out + d); } for (; d < size; ++d) { diff --git a/sgl-kernel/csrc/cpu/mamba/conv.cpp b/sgl-kernel/csrc/cpu/mamba/conv.cpp index a312b44e20da..425d839df226 100644 --- a/sgl-kernel/csrc/cpu/mamba/conv.cpp +++ b/sgl-kernel/csrc/cpu/mamba/conv.cpp @@ -177,14 +177,13 @@ struct tinygemm_kernel { using fVec = at::vec::Vectorized; using bVec = at::vec::Vectorized; - const fVec one = fVec(1.f); auto storec = [&](auto i, int64_t m) { constexpr int col = i; fVec x0 = fVec(vc[col * 2 + 0]); fVec x1 = fVec(vc[col * 2 + 1]); if constexpr (has_silu) { - x0 = x0 / (one + x0.neg().exp_u20()); - x1 = x1 / (one + x1.neg().exp_u20()); + x0 = fast_silu(x0); + x1 = fast_silu(x1); } bVec out_vec = convert_from_float_ext(x0, x1); out_vec.store(C + m * lda + col * 32); diff --git a/sgl-kernel/csrc/cpu/mamba/fla.cpp b/sgl-kernel/csrc/cpu/mamba/fla.cpp index 2abfa25fa4c7..641d7056b5c4 100644 --- a/sgl-kernel/csrc/cpu/mamba/fla.cpp +++ b/sgl-kernel/csrc/cpu/mamba/fla.cpp @@ -1146,14 +1146,10 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl( int64_t d; #pragma GCC unroll 4 for (d = 0; d <= head_dim - VecSize; d += VecSize) { - bVec q_bvec = bVec::loadu(q_ptr + q_offset + d); - fVec q_fvec0, q_fvec1; - std::tie(q_fvec0, q_fvec1) = at::vec::convert_to_float(q_bvec); + auto [q_fvec0, q_fvec1] = load_float_vec2(q_ptr + q_offset + d); sum_q_fvec += q_fvec0 * q_fvec0; sum_q_fvec += q_fvec1 * q_fvec1; - bVec k_bvec = bVec::loadu(k_ptr + k_offset + d); - fVec k_fvec0, k_fvec1; - std::tie(k_fvec0, k_fvec1) = at::vec::convert_to_float(k_bvec); + auto [k_fvec0, k_fvec1] = load_float_vec2(k_ptr + k_offset + d); sum_k_fvec += k_fvec0 * k_fvec0; sum_k_fvec += k_fvec1 * k_fvec1; } @@ -1200,14 +1196,11 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl( fVec kv_mem_vec1 = fVec(float(0)); for (int di = 0; di < head_dim; ++di) { fVec k_val_vec = fVec(k_ptr[k_offset + di] * k_scale); - fVec state_vec0 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi); - fVec state_vec1 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi + fVecSize); + auto [state_vec0, state_vec1] = load_float_vec2(state_ptr + state_offset + di * v_head_dim + dvi); kv_mem_vec0 = kv_mem_vec0 + state_vec0 * g_val_exp_vec * k_val_vec; kv_mem_vec1 = kv_mem_vec1 + state_vec1 * g_val_exp_vec * k_val_vec; } - bVec v_bvec = bVec::loadu(v_ptr + v_offset + dvi); - fVec v_vec0, v_vec1; - std::tie(v_vec0, v_vec1) = at::vec::convert_to_float(v_bvec); + auto [v_vec0, v_vec1] = load_float_vec2(v_ptr + v_offset + dvi); fVec dt_vec0 = (v_vec0 - kv_mem_vec0) * beta_vec; fVec dt_vec1 = (v_vec1 - kv_mem_vec1) * beta_vec; fVec o_vec0 = fVec(float(0)); @@ -1215,8 +1208,7 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl( for (int di = 0; di < head_dim; ++di) { fVec q_vec = fVec(q_ptr[q_offset + di] * q_scale); fVec k_vec = fVec(k_ptr[k_offset + di] * k_scale); - fVec state_vec0 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi); - fVec state_vec1 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi + fVecSize); + auto [state_vec0, state_vec1] = load_float_vec2(state_ptr + state_offset + di * v_head_dim + dvi); state_vec0 = state_vec0 * g_val_exp_vec + k_vec * dt_vec0; state_vec1 = state_vec1 * g_val_exp_vec + k_vec * dt_vec1; o_vec0 = o_vec0 + state_vec0 * q_vec * scale_vec; @@ -1270,20 +1262,15 @@ void fused_gdn_gating_kernel_impl( for (int64_t i = begin; i < end; ++i) { int64_t j = 0; for (; j < num_heads - (num_heads % vec_size); j += vec_size) { - fVec A_log_vec0 = fVec::loadu(A_log + j); - fVec A_log_vec1 = fVec::loadu(A_log + j + fvec_size); - bVec dt_bias_vec = bVec::loadu(dt_bias + j); - bVec a_bvec = bVec::loadu(a + i * num_heads + j); - bVec b_bvec = bVec::loadu(b + i * num_heads + j); - fVec a0, a1, dt_bias_vec0, dt_bias_vec1, b0, b1; - std::tie(a0, a1) = at::vec::convert_to_float(a_bvec); - std::tie(b0, b1) = at::vec::convert_to_float(b_bvec); - std::tie(dt_bias_vec0, dt_bias_vec1) = at::vec::convert_to_float(dt_bias_vec); + auto [A_log_vec0, A_log_vec1] = load_float_vec2(A_log + j); + auto [dt_bias_vec0, dt_bias_vec1] = load_float_vec2(dt_bias + j); + auto [a0, a1] = load_float_vec2(a + i * num_heads + j); + auto [b0, b1] = load_float_vec2(b + i * num_heads + j); fVec g0 = neg_one * A_log_vec0.exp_u20() * softplus(a0 + dt_bias_vec0); fVec g1 = neg_one * A_log_vec1.exp_u20() * softplus(a1 + dt_bias_vec1); - fVec beta0 = one / (one + (neg_one * b0).exp_u20()); - fVec beta1 = one / (one + (neg_one * b1).exp_u20()); + fVec beta0 = fast_sigmoid(b0); + fVec beta1 = fast_sigmoid(b1); g0.store(out + i * num_heads + j); g1.store(out + i * num_heads + j + fvec_size); @@ -1318,21 +1305,15 @@ void fused_gdn_gating_kernel_impl( for (int64_t i = begin; i < end; ++i) { int64_t j = 0; for (; j < num_heads - (num_heads % vec_size); j += vec_size) { - bVec A_log_bvec = bVec::loadu(A_log + j); - fVec A_log_vec0, A_log_vec1; - std::tie(A_log_vec0, A_log_vec1) = at::vec::convert_to_float(A_log_bvec); - bVec dt_bias_vec = bVec::loadu(dt_bias + j); - bVec a_bvec = bVec::loadu(a + i * num_heads + j); - bVec b_bvec = bVec::loadu(b + i * num_heads + j); - fVec a0, a1, dt_bias_vec0, dt_bias_vec1, b0, b1; - std::tie(a0, a1) = at::vec::convert_to_float(a_bvec); - std::tie(b0, b1) = at::vec::convert_to_float(b_bvec); - std::tie(dt_bias_vec0, dt_bias_vec1) = at::vec::convert_to_float(dt_bias_vec); + auto [A_log_vec0, A_log_vec1] = load_float_vec2(A_log + j); + auto [dt_bias_vec0, dt_bias_vec1] = load_float_vec2(dt_bias + j); + auto [a0, a1] = load_float_vec2(a + i * num_heads + j); + auto [b0, b1] = load_float_vec2(b + i * num_heads + j); fVec g0 = neg_one * A_log_vec0.exp_u20() * softplus(a0 + dt_bias_vec0); fVec g1 = neg_one * A_log_vec1.exp_u20() * softplus(a1 + dt_bias_vec1); - fVec beta0 = one / (one + (neg_one * b0).exp_u20()); - fVec beta1 = one / (one + (neg_one * b1).exp_u20()); + fVec beta0 = fast_sigmoid(b0); + fVec beta1 = fast_sigmoid(b1); g0.store(out + i * num_heads + j); g1.store(out + i * num_heads + j + fvec_size); diff --git a/sgl-kernel/csrc/cpu/moe.cpp b/sgl-kernel/csrc/cpu/moe.cpp index 797a6edc7c2b..9b2ed699cf86 100644 --- a/sgl-kernel/csrc/cpu/moe.cpp +++ b/sgl-kernel/csrc/cpu/moe.cpp @@ -118,96 +118,6 @@ int moe_align_block_size( return num_tokens_post_pad; } -// silu : shape leading dimension -// input0 [m_size, BLOCK_N] BLOCK_N -// input1 [m_size, BLOCK_N] BLOCK_N -// output [M * topk, N] N -template -inline void silu_and_mul( - scalar_t* __restrict__ output, - const float* __restrict__ input0, // x: x0, x1 - const float* __restrict__ input1, // y: y0, y1 - int64_t m_size, - int64_t N) { - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - - const fVec one = fVec(1.f); - - // no remainder - for (int64_t m = 0; m < m_size; ++m) { - scalar_t* __restrict__ out = output + m * N; - const float* __restrict__ x = input0 + m * BLOCK_N; - const float* __restrict__ y = input1 + m * BLOCK_N; - - for (int64_t d = 0; d < BLOCK_N; d += bVec::size()) { - fVec x0 = fVec::loadu(x + d); - fVec x1 = fVec::loadu(x + d + fVec::size()); - fVec y0 = fVec::loadu(y + d); - fVec y1 = fVec::loadu(y + d + fVec::size()); - // silu - x0 = x0 / (one + x0.neg().exp_u20()); - x1 = x1 / (one + x1.neg().exp_u20()); - // mul - x0 = x0 * y0; - x1 = x1 * y1; - // convert - bVec out_vec = convert_from_float_ext(x0, x1); - out_vec.store(out + d); - } - } -} - -template -inline void clamp_sigmoid_and_mul( - scalar_t* __restrict__ output, - const float* __restrict__ input0, - int64_t m_size, - int64_t N, - const float alpha, - const float limit, - int64_t offset) { - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - - const fVec one = fVec(1.f); - const fVec zero = fVec(0.f); - const fVec limit_v = fVec(limit); - const fVec nlimit_v = fVec(-limit); - const fVec alpha_v = fVec(alpha); - - // no remainder - for (int64_t m = 0; m < m_size; ++m) { - scalar_t* __restrict__ out = output + m * N; - const float* __restrict__ cur_ptr = input0 + m * BLOCK_N; - for (int64_t d = 0; d < BLOCK_N; d += bVec::size()) { - float tmp_glu0[fVec::size()]; // 16 - float tmp_linear0[fVec::size()]; // 16 - - // interleaved: x[2i] = glu, x[2i+1] = linear - for (int j = 0; j < fVec::size(); ++j) { - // x0 [0,2,..30] - tmp_glu0[j] = cur_ptr[d + j * 2]; - // y0 [1,3,...31] - tmp_linear0[j] = cur_ptr[d + j * 2 + 1]; - } - fVec x0 = fVec::loadu(tmp_glu0); - fVec y0 = fVec::loadu(tmp_linear0); - - // clamp - x0 = at::vec::minimum(x0, limit_v); - y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0)); - // x * sigmoid(x * alpha) - x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20()); - // (y + 1) * x - y0 = y0 + one; - x0 = x0 * y0; - // // convert - convert_from_float_and_store(out + d / 2 + offset, x0); - } - } -} - template struct tinygemm_kernel_nn2 { static inline void apply( @@ -284,23 +194,17 @@ struct tinygemm_kernel_nn2 { Unroll{}(compute, k); } - using Vec = at::vec::Vectorized; - const Vec one = Vec(1.f); auto storec = [&](auto i) { constexpr int row = i / COLS; constexpr int col = i % COLS; // for COLS = 2, 4 use 512bit store if constexpr (col % 2 == 0) { - Vec x0 = vc0[row * COLS + col + 0]; - Vec x1 = vc0[row * COLS + col + 1]; - Vec y0 = vc1[row * COLS + col + 0]; - Vec y1 = vc1[row * COLS + col + 1]; - // silu - x0 = x0 / (one + x0.neg().exp_u20()); - x1 = x1 / (one + x1.neg().exp_u20()); - // mul - x0 = x0 * y0; - x1 = x1 * y1; + __m512 x0 = vc0[row * COLS + col + 0]; + __m512 x1 = vc0[row * COLS + col + 1]; + __m512 y0 = vc1[row * COLS + col + 0]; + __m512 y1 = vc1[row * COLS + col + 1]; + x0 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x0), y0); + x1 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x1), y1); _mm512_storeu_si512( reinterpret_cast<__m512i*>((C + row * ldc + col * 16)), @@ -638,11 +542,15 @@ void fused_experts_kernel_impl( // 1.d silu and mul const int64_t offset = offsets[mb]; if (act_func == CPUActMethod::silu_and_mul && use_brgemm) { - silu_and_mul(ic1 + offset * N + nb * BLOCK_N, C0, C1, m_size, N); + for (int64_t m = 0; m < m_size; ++m) { + silu_and_mul_stub(ic1 + (offset + m) * N + nb * BLOCK_N, C0 + m * BLOCK_N, C1 + m * BLOCK_N, BLOCK_N); + } } else if (act_func == CPUActMethod::swiglu) { - clamp_sigmoid_and_mul(ic1 + offset * N, C0, m_size, N, alpha, limit, 0 + nb * BLOCK_N / 2); - clamp_sigmoid_and_mul( - ic1 + offset * N, C1, m_size, N, alpha, limit, N / 2 + nb * BLOCK_N / 2); + for (int64_t m = 0; m < m_size; ++m) { + scalar_t* __restrict__ ic1_row = ic1 + (offset + m) * N; + clamp_sigmoid_and_mul_stub(ic1_row + nb * BLOCK_N / 2, C0 + m * BLOCK_N, BLOCK_N / 2, alpha, limit); + clamp_sigmoid_and_mul_stub(ic1_row + N / 2 + nb * BLOCK_N / 2, C1 + m * BLOCK_N, BLOCK_N / 2, alpha, limit); + } } }); @@ -811,7 +719,9 @@ void shared_expert_kernel_impl( /* C */ C1); // 1.d silu and mul - silu_and_mul(ic1 + mb * BLOCK_M * N + nb * BLOCK_N, C0, C1, m_size, N); + for (int64_t m = 0; m < m_size; ++m) { + silu_and_mul_stub(ic1 + (mb * BLOCK_M + m) * N + nb * BLOCK_N, C0 + m * BLOCK_N, C1 + m * BLOCK_N, BLOCK_N); + } } else { // fused 1.bcd: silu_and_mul(A @ B0, A @ B1) tinygemm_kernel( diff --git a/sgl-kernel/csrc/cpu/moe.h b/sgl-kernel/csrc/cpu/moe.h index 71f829c26425..71b28d868ee2 100644 --- a/sgl-kernel/csrc/cpu/moe.h +++ b/sgl-kernel/csrc/cpu/moe.h @@ -59,9 +59,7 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const input_t* __restrict_ #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { auto [x0, x1] = load_float_vec2(input + d); - x0 = x0 * weight_vec; - x1 = x1 * weight_vec; - bVec out_vec = convert_from_float_ext(x0, x1); + bVec out_vec = convert_from_float_ext(x0 * weight_vec, x1 * weight_vec); out_vec.store(out + d); } for (; d < size; ++d) { @@ -86,10 +84,7 @@ inline void sum_stub(scalar_t* __restrict__ out, const scalar_t* __restrict__ in fVec sum_fvec0 = fVec(0.f); fVec sum_fvec1 = fVec(0.f); for (int t = 0; t < topk; ++t) { - bVec x_bvec = bVec::loadu(input + t * K + d); - fVec x_fvec0, x_fvec1; - std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); - + auto [x_fvec0, x_fvec1] = load_float_vec2(input + t * K + d); sum_fvec0 += x_fvec0; sum_fvec1 += x_fvec1; } @@ -132,11 +127,7 @@ inline void add_mul_stub( #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { auto [x0, x1] = load_float_vec2(input + d); - - bVec y_bvec = bVec::loadu(input2 + d); - fVec y0, y1; - std::tie(y0, y1) = at::vec::convert_to_float(y_bvec); - + auto [y0, y1] = load_float_vec2(input2 + d); x0 = x0 + y0 * s_vec; x1 = x1 + y1 * s_vec; bVec out_vec = convert_from_float_ext(x0, x1); @@ -147,31 +138,52 @@ inline void add_mul_stub( } } -template +template inline void silu_and_mul_stub( - scalar_t* __restrict__ out, const scalar_t* __restrict__ input, const scalar_t* __restrict__ input2, int64_t size) { + scalar_t* __restrict__ out, const input_t* __restrict__ input, const input_t* __restrict__ input2, int64_t size) { + static_assert( + std::is_same_v || std::is_same_v, + "silu_and_mul_stub only supports input_t == float or input_t == scalar_t"); using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; - const fVec one = fVec(1.f); // no remainder #pragma GCC unroll 4 for (int64_t d = 0; d < size; d += bVec::size()) { - bVec x = bVec::loadu(input + d); - fVec x0, x1; - std::tie(x0, x1) = at::vec::convert_to_float(x); - bVec y = bVec::loadu(input2 + d); - fVec y0, y1; - std::tie(y0, y1) = at::vec::convert_to_float(y); - x0 = x0 / (one + x0.neg().exp_u20()); - x1 = x1 / (one + x1.neg().exp_u20()); - x0 = x0 * y0; - x1 = x1 * y1; + auto [x0, x1] = load_float_vec2(input + d); + auto [y0, y1] = load_float_vec2(input2 + d); + x0 = fast_silu(x0) * y0; + x1 = fast_silu(x1) * y1; bVec out_vec = convert_from_float_ext(x0, x1); out_vec.store(out + d); } } +template +inline void clamp_sigmoid_and_mul_stub( + scalar_t* __restrict__ out, const input_t* __restrict__ input, int64_t size, const float alpha, const float limit) { + static_assert( + std::is_same_v || std::is_same_v, + "clamp_sigmoid_and_mul_stub only supports input_t == float or input_t == scalar_t"); + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + const fVec one = fVec(1.f); + const fVec limit_v = fVec(limit); + const fVec nlimit_v = fVec(-limit); + const fVec alpha_v = fVec(alpha); + +#pragma GCC unroll 4 + for (int64_t d = 0; d < 2 * size; d += bVec::size()) { + auto [x0_, y0_] = load_float_vec2(input + d); + auto [x0, y0] = at::vec::deinterleave2(x0_, y0_); + + x0 = at::vec::minimum(x0, limit_v); + y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0)); + x0 = fast_sigmoid_glu(x0, alpha_v) * (y0 + one); + store_from_float_ext(out + d / 2, x0); + } +} + template inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ input, float weight, int64_t size) { using bVec = at::vec::Vectorized; @@ -181,9 +193,8 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec data0 = fVec::loadu(input + d) * weight_vec; - fVec data1 = fVec::loadu(input + d + fVec::size()) * weight_vec; - bVec out_vec = convert_from_float_ext(data0, data1); + auto [x0, x1] = load_float_vec2(input + d); + bVec out_vec = convert_from_float_ext(x0 * weight_vec, x1 * weight_vec); out_vec.store(out + d); } for (; d < size; ++d) { @@ -198,8 +209,7 @@ inline void add_bias_stub(float* __restrict__ input, const float* __restrict__ i int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - fVec x_fvec = fVec::loadu(input + d); - fVec y_fvec = fVec::loadu(input2 + d); + auto [x_fvec, y_fvec] = load_float_vec2(input + d); x_fvec = x_fvec + y_fvec; x_fvec.store(input + d); } @@ -217,63 +227,11 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const scalar_t* __restrict int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - bVec x = bVec::loadu(input + d); - fVec x0, x1; - std::tie(x0, x1) = at::vec::convert_to_float(x); - x0 = x0 * weight_vec; - x1 = x1 * weight_vec; - bVec out_vec = convert_from_float_ext(x0, x1); + auto [x0, x1] = load_float_vec2(input + d); + bVec out_vec = convert_from_float_ext(x0 * weight_vec, x1 * weight_vec); out_vec.store(out + d); } for (; d < size; ++d) { out[d] = static_cast(input[d] * weight); } } - -template -inline void clamp_sigmoid_and_mul_stub( - scalar_t* __restrict__ out, - const scalar_t* __restrict__ input, - int64_t size, - const float alpha, - const float limit) { - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - const fVec one = fVec(1.f); - const fVec zero = fVec(0.f); - const fVec limit_v = fVec(limit); - const fVec nlimit_v = fVec(-limit); - const fVec alpha_v = fVec(alpha); - - // no remainder -#pragma GCC unroll 4 - for (int64_t d = 0; d < size; d += bVec::size()) { - bVec x = bVec::loadu(input + d); - fVec x0_, y0_; - std::tie(x0_, y0_) = at::vec::convert_to_float(x); - float tmp_buffer[fVec::size() * 2]; // 32 - float tmp_glu[fVec::size()]; // 16 - float tmp_linear[fVec::size()]; // 16 - x0_.store(tmp_buffer); - y0_.store(tmp_buffer + fVec::size()); - // interleaved: x[2i] = glu, x[2i+1] = linear - for (int j = 0; j < fVec::size(); ++j) { - // x0 [0,2,..30] - tmp_glu[j] = tmp_buffer[j * 2]; - // y0 [1,3,...31] - tmp_linear[j] = tmp_buffer[j * 2 + 1]; - } - fVec x0 = fVec::loadu(tmp_glu); - fVec y0 = fVec::loadu(tmp_linear); - - // clamp - x0 = at::vec::minimum(x0, limit_v); - y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0)); - // x * sigmoid(x * alpha) - x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20()); - // (y + 1) * x - y0 = y0 + one; - x0 = x0 * y0; - convert_from_float_and_store(out + d / 2, x0); - } -} diff --git a/sgl-kernel/csrc/cpu/moe_int8.cpp b/sgl-kernel/csrc/cpu/moe_int8.cpp index b147b65f60fe..b5b9f6aecc28 100644 --- a/sgl-kernel/csrc/cpu/moe_int8.cpp +++ b/sgl-kernel/csrc/cpu/moe_int8.cpp @@ -48,16 +48,14 @@ inline void silu_and_mul( vc1[col] = _mm512_mul_ps(_mm512_mul_ps(vc1[col], vas), vbs1[col]); }; - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - const fVec one = fVec(1.f); auto silu_and_mul = [&](auto col) { - fVec x = fVec(vc0[col]); - fVec y = fVec(vc1[col]); - x = x / (one + x.neg().exp_u20()); - vc0[col] = x * y; + __m512 x = vc0[col]; + __m512 y = vc1[col]; + vc0[col] = _mm512_mul_ps(_mm512_rcp14_silu_ps(x), y); }; + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; auto storec = [&](auto col, int64_t m) { if constexpr (col % 2 == 0) { fVec x0 = fVec(vc0[col + 0]); @@ -224,23 +222,17 @@ struct tinygemm_kernel_vnni { }; Unroll{}(scalec); - using Vec = at::vec::Vectorized; - const Vec one = Vec(1.f); auto storec = [&](auto i) { constexpr int row = i / COLS; constexpr int col = i % COLS; // for COLS = 2, 4 use 512bit store if constexpr (col % 2 == 0) { - Vec x0 = _mm512_castsi512_ps(vc0[row * COLS + col + 0]); - Vec x1 = _mm512_castsi512_ps(vc0[row * COLS + col + 1]); - Vec y0 = _mm512_castsi512_ps(vc1[row * COLS + col + 0]); - Vec y1 = _mm512_castsi512_ps(vc1[row * COLS + col + 1]); - // silu - x0 = x0 / (one + x0.neg().exp_u20()); - x1 = x1 / (one + x1.neg().exp_u20()); - // mul - x0 = x0 * y0; - x1 = x1 * y1; + __m512 x0 = _mm512_castsi512_ps(vc0[row * COLS + col + 0]); + __m512 x1 = _mm512_castsi512_ps(vc0[row * COLS + col + 1]); + __m512 y0 = _mm512_castsi512_ps(vc1[row * COLS + col + 0]); + __m512 y1 = _mm512_castsi512_ps(vc1[row * COLS + col + 1]); + x0 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x0), y0); + x1 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x1), y1); _mm512_storeu_si512( reinterpret_cast<__m512i*>((C + row * ldc + col * 16)), diff --git a/sgl-kernel/csrc/cpu/qkv_proj.cpp b/sgl-kernel/csrc/cpu/qkv_proj.cpp index f929efb92c42..f0489e4fe619 100644 --- a/sgl-kernel/csrc/cpu/qkv_proj.cpp +++ b/sgl-kernel/csrc/cpu/qkv_proj.cpp @@ -240,9 +240,7 @@ inline float reduce(const scalar_t* __restrict__ x, int64_t size) { // no remainder #pragma GCC unroll 4 for (int64_t d = 0; d < size; d += bVec::size()) { - bVec x_bvec = bVec::loadu(x + d); - fVec x_fvec0, x_fvec1; - std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); + auto [x_fvec0, x_fvec1] = load_float_vec2(x + d); sum_fvec += x_fvec0 * x_fvec0; sum_fvec += x_fvec1 * x_fvec1; } @@ -259,12 +257,8 @@ inline void map2(scalar_t* y, const scalar_t* x, const scalar_t* __restrict__ w, // no remainder #pragma GCC unroll 4 for (int64_t d = 0; d < size; d += bVec::size()) { - bVec x_bvec = bVec::loadu(x + d); - fVec x_fvec0, x_fvec1; - std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); - bVec w_bvec = bVec::loadu(w + d); - fVec w_fvec0, w_fvec1; - std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec); + auto [x_fvec0, x_fvec1] = load_float_vec2(x + d); + auto [w_fvec0, w_fvec1] = load_float_vec2(w + d); x_fvec0 = x_fvec0 * scale_fvec * w_fvec0; x_fvec1 = x_fvec1 * scale_fvec * w_fvec1; bVec out_bvec = convert_from_float_ext(x_fvec0, x_fvec1); diff --git a/sgl-kernel/csrc/cpu/vec.h b/sgl-kernel/csrc/cpu/vec.h index 22c02a9aeb64..3d59aa064af9 100644 --- a/sgl-kernel/csrc/cpu/vec.h +++ b/sgl-kernel/csrc/cpu/vec.h @@ -17,11 +17,11 @@ inline Vectorized convert_from_float_ext(const Vectorized& a, c } template -inline void convert_from_float_and_store(scalar_t* out, const Vectorized& a) { - float out_buffer[at::vec::Vectorized::size()]; +inline void store_from_float_ext(scalar_t* out, const Vectorized& a) { + float out_buffer[Vectorized::size()]; a.store(out_buffer); - for (int i = 0; i < 16; i++) { - out[i] = (scalar_t)out_buffer[i]; + for (int i = 0; i < Vectorized::size(); ++i) { + out[i] = static_cast(out_buffer[i]); } } @@ -55,8 +55,14 @@ convert_from_float_ext(const Vectorized& a, const Vectorize } template <> -inline void convert_from_float_and_store(at::BFloat16* out, const Vectorized& a) { - _mm256_storeu_si256((__m256i*)out, (__m256i)(_mm512_cvtneps_pbh(__m512(a)))); +inline void store_from_float_ext(at::BFloat16* out, const Vectorized& a) { + _mm256_storeu_si256(reinterpret_cast<__m256i*>(out), (__m256i)(_mm512_cvtneps_pbh(__m512(a)))); +} + +template <> +inline void store_from_float_ext(at::Half* out, const Vectorized& a) { + _mm256_storeu_si256( + reinterpret_cast<__m256i*>(out), _mm512_cvtps_ph(__m512(a), _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC)); } #define CVT_BF16_TO_FP32(a) _mm512_castsi512_ps(_mm512_slli_epi32(_mm512_cvtepu16_epi32(a), 16)) @@ -557,6 +563,51 @@ inline __attribute__((always_inline)) __m512 _mm512_fexp_u20_ps(const __m512 val // final interpretation to float return _mm512_castsi512_ps(casted_integer); } + +// sigmoid(x) = 1 / (1 + exp(-x)); avoid vdivps via rcp14 +inline __attribute__((always_inline)) __m512 _mm512_rcp14_sigmoid_ps(__m512 x) { + __m512 minus_x = _mm512_xor_ps(_mm512_set1_ps(-0.f), x); + __m512 denom = _mm512_add_ps(_mm512_exp_u20_ps(minus_x), _mm512_set1_ps(1.f)); + return _mm512_rcp14_ps(denom); +} + +// SiLU(x) = x * sigmoid(x) +inline __attribute__((always_inline)) __m512 _mm512_rcp14_silu_ps(__m512 x) { + return _mm512_mul_ps(x, _mm512_rcp14_sigmoid_ps(x)); +} + +// x * sigmoid(x * alpha) for clamped SwiGLU +inline __attribute__((always_inline)) __m512 _mm512_rcp14_sigmoid_glu_ps(__m512 x, __m512 alpha) { + __m512 xa = _mm512_mul_ps(x, alpha); + return _mm512_mul_ps(x, _mm512_rcp14_sigmoid_ps(xa)); +} + #endif +inline at::vec::Vectorized fast_sigmoid(const at::vec::Vectorized& x) { +#if defined(CPU_CAPABILITY_AVX512) + return at::vec::Vectorized(_mm512_rcp14_sigmoid_ps(x)); +#else + const auto one = at::vec::Vectorized(1.f); + return one / (one + x.neg().exp_u20()); +#endif +} + +inline at::vec::Vectorized fast_silu(const at::vec::Vectorized& x) { +#if defined(CPU_CAPABILITY_AVX512) + return at::vec::Vectorized(_mm512_rcp14_silu_ps(x)); +#else + return x * fast_sigmoid(x); +#endif +} + +inline at::vec::Vectorized +fast_sigmoid_glu(const at::vec::Vectorized& x, const at::vec::Vectorized& alpha) { +#if defined(CPU_CAPABILITY_AVX512) + return at::vec::Vectorized(_mm512_rcp14_sigmoid_glu_ps(x, alpha)); +#else + return x * fast_sigmoid(x * alpha); +#endif +} + } // anonymous namespace From 92fcf42fde73e6cbb5e73140c73ea3b2f989fdd5 Mon Sep 17 00:00:00 2001 From: mingfeima Date: Thu, 16 Jul 2026 15:15:09 +0000 Subject: [PATCH 2/5] fix error --- sgl-kernel/csrc/cpu/moe.h | 3 +- sgl-kernel/csrc/cpu/moe_fp8.cpp | 4 +- test/registered/cpu/test_moe.py | 74 +++++++++++++++++++++++++++++++++ 3 files changed, 78 insertions(+), 3 deletions(-) diff --git a/sgl-kernel/csrc/cpu/moe.h b/sgl-kernel/csrc/cpu/moe.h index 71b28d868ee2..275a06b39e50 100644 --- a/sgl-kernel/csrc/cpu/moe.h +++ b/sgl-kernel/csrc/cpu/moe.h @@ -209,7 +209,8 @@ inline void add_bias_stub(float* __restrict__ input, const float* __restrict__ i int64_t d; #pragma GCC unroll 4 for (d = 0; d <= size - kVecSize; d += kVecSize) { - auto [x_fvec, y_fvec] = load_float_vec2(input + d); + fVec x_fvec = fVec::loadu(input + d); + fVec y_fvec = fVec::loadu(input2 + d); x_fvec = x_fvec + y_fvec; x_fvec.store(input + d); } diff --git a/sgl-kernel/csrc/cpu/moe_fp8.cpp b/sgl-kernel/csrc/cpu/moe_fp8.cpp index 53a643d47972..9dd2bb0ee7bc 100644 --- a/sgl-kernel/csrc/cpu/moe_fp8.cpp +++ b/sgl-kernel/csrc/cpu/moe_fp8.cpp @@ -126,8 +126,8 @@ void fused_experts_fp_kernel_impl( } else if (act_func == CPUActMethod::swiglu) { at::parallel_for(0, M * topk, 0, [&](int64_t begin, int64_t end) { for (int64_t m = begin; m < end; ++m) { - clamp_sigmoid_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, N, alpha, limit); - clamp_sigmoid_and_mul_stub(ic1 + m * N + N / 2, ic0 + m * 2 * N + N, N, alpha, limit); + clamp_sigmoid_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, N / 2, alpha, limit); + clamp_sigmoid_and_mul_stub(ic1 + m * N + N / 2, ic0 + m * 2 * N + N, N / 2, alpha, limit); } }); } diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index 0ede6aa40c0a..f20d03865f41 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -251,6 +251,80 @@ def test_fp8_moe(self, M, N, K, E, topk): atol = rtol = precision[dtype] torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) + @parametrize( + m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] + ) + def test_fp8_moe_bias(self, m, n, k, e, topk, renormalize): + dtype = torch.bfloat16 + + a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 + + w1_fp32 = torch.randn((e, 2 * n, k), device="cpu", dtype=torch.float32) / 10 + w1 = (w1_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) + w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float32) / 10 + + w2_fp32 = torch.randn((e, k, n), device="cpu", dtype=torch.float32) / 10 + w2 = (w2_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) + w2_b = torch.randn((e, k), device="cpu", dtype=torch.float32) / 10 + + w1s = ( + torch.randn(e, math.ceil(2 * n / BLOCK_N), math.ceil(k / BLOCK_K)) + * factor_for_scale + ) + w2s = ( + torch.randn(e, math.ceil(k / BLOCK_N), math.ceil(n / BLOCK_K)) + * factor_for_scale + ) + + w1_scaled = scaled_weight(w1, w1s).to(dtype) + w2_scaled = scaled_weight(w2, w2s).to(dtype) + + score = torch.randn((m, e), device="cpu", dtype=dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) + alpha = 1.702 + limit = 7.0 + + ref_out = torch_naive_fused_moe_gptoss( + a, + w1_scaled, + w2_scaled, + w1_b, + w2_b, + topk_weight, + topk_ids, + renormalize, + alpha, + limit, + e, + ) + + w1 = kernel.convert_weight_packed(w1) + w2 = kernel.convert_weight_packed(w2) + + out = kernel.fused_experts_cpu( + a, + w1, + w2, + topk_weight, + topk_ids.to(torch.int32), + False, + CPUQuantMethod.FP8_W8A16, + w1s, + w2s, + None, + None, + [BLOCK_N, BLOCK_K], + w1_b, + w2_b, + alpha, + limit, + True, + ) + + atol = rtol = precision[dtype] + torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) + @parametrize(M=[2, 121], N=[352, 512], K=[256, 320], E=[8], topk=[4]) def test_mxfp4_moe(self, M, N, K, E, topk): dtype = torch.bfloat16 From 8c2ac79abd9b21e3c992bbc96956b23a8721b76b Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 16 Jul 2026 07:53:07 +0000 Subject: [PATCH 3/5] =?UTF-8?q?Refactor=20CPU=20MoE=20tests:=20merge=20?= =?UTF-8?q?=C2=B1bias=20cases=20via=20bias=3D[False,=20True]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add run_fused_experts / make_routing helpers to drop None arg walls - Merge bf16/fp8/mxfp4 non-bias and bias tests into one parametrized case - bias=True also covers clamped SwiGLU (alpha/limit); bias=False keeps silu Co-authored-by: Ma Mingfei --- test/registered/cpu/test_moe.py | 587 +++++++++++++------------------- 1 file changed, 231 insertions(+), 356 deletions(-) diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index f20d03865f41..39f0e70f2e51 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -32,303 +32,233 @@ register_cpu_ci(est_time=10, suite="base-b-test-cpu") - -def fused_moe(a, w1, w2, score, topk, renormalize, prepack): - - G = 1 - topk_group = 1 - - B, D = a.shape - topk_weights = torch.empty(B, topk, dtype=torch.float32) - topk_ids = torch.empty(B, topk, dtype=torch.int32) - topk_weights, topk_ids = kernel.grouped_topk_cpu( - a, score, topk, renormalize, G, topk_group, 0, None, None - ) - - packed_w1 = kernel.convert_weight_packed(w1) if prepack else w1 - packed_w2 = kernel.convert_weight_packed(w2) if prepack else w2 - - inplace = True +# GPT-OSS / MiniMax-style SwiGLU clamp params (used when bias=True) +_SWIGLU_ALPHA = 1.702 +_SWIGLU_LIMIT = 7.0 + + +def run_fused_experts( + a, + w1, + w2, + topk_weight, + topk_ids, + *, + quant=CPUQuantMethod.UNQUANT, + w1_scale=None, + w2_scale=None, + w1_zp=None, + w2_zp=None, + block_size=None, + w1_bias=None, + w2_bias=None, + alpha=None, + limit=None, + is_vnni=True, + inplace=False, +): return kernel.fused_experts_cpu( a, - packed_w1, - packed_w2, - topk_weights, - topk_ids, + w1, + w2, + topk_weight, + topk_ids.to(torch.int32), inplace, - CPUQuantMethod.UNQUANT, - None, - None, - None, - None, - None, - None, - None, - None, - None, - prepack, + quant, + w1_scale, + w2_scale, + w1_zp, + w2_zp, + block_size, + w1_bias, + w2_bias, + alpha, + limit, + is_vnni, ) -class TestFusedExperts(CustomTestCase): +def make_routing(m, e, topk, dtype, renormalize=False): + score = torch.randn((m, e), dtype=dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) + if renormalize: + topk_weight = topk_weight / topk_weight.sum(dim=-1, keepdim=True) + return topk_weight, topk_ids - @parametrize(m=[2, 114], n=[32], k=[32], e=[4], topk=[2], renormalize=[False, True]) - def test_bf16_moe(self, m, n, k, e, topk, renormalize): - dtype = torch.bfloat16 - prepack = True - a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 - w1 = torch.randn((e, 2 * n, k), device="cpu", dtype=dtype) / 10 - w2 = torch.randn((e, k, n), device="cpu", dtype=dtype) / 10 - score = torch.randn((m, e), device="cpu", dtype=dtype) - - torch_output = torch_naive_fused_moe(a, w1, w2, score, topk, renormalize) - fused_output = fused_moe(a, w1, w2, score, topk, renormalize, prepack) - - atol = rtol = precision[torch_output.dtype] - torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) +class TestFusedExperts(CustomTestCase): @parametrize( - m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] + m=[2, 32], + n=[32, 128], + k=[32, 128], + e=[4], + topk=[2], + renormalize=[False, True], + bias=[False, True], ) - def test_bf16_moe_bias(self, m, n, k, e, topk, renormalize): + def test_bf16_moe(self, m, n, k, e, topk, renormalize, bias): dtype = torch.bfloat16 - - a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 - w1 = torch.randn((e, 2 * n, k), device="cpu", dtype=dtype) / 10 - w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float) / 10 - w2 = torch.randn((e, k, n), device="cpu", dtype=dtype) / 10 - w2_b = torch.randn((e, k), device="cpu", dtype=torch.float) / 10 - score = torch.randn((m, e), device="cpu", dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) - alpha = 1.702 - limit = 7.0 - torch_output = torch_naive_fused_moe_gptoss( - a, w1, w2, w1_b, w2_b, topk_weight, topk_ids, renormalize, alpha, limit, e - ) + a = torch.randn((m, k), dtype=dtype) / 10 + w1 = torch.randn((e, 2 * n, k), dtype=dtype) / 10 + w2 = torch.randn((e, k, n), dtype=dtype) / 10 packed_w1 = kernel.convert_weight_packed(w1) packed_w2 = kernel.convert_weight_packed(w2) - fused_output = torch.ops.sgl_kernel.fused_experts_cpu( - a, - packed_w1, - packed_w2, - topk_weight, - topk_ids.to(torch.int), - False, # inplace # See [Note] inplace should be False in fused_experts. - CPUQuantMethod.UNQUANT, - None, # w1_scale - None, # w2_scale - None, # w1_zp - None, # w2_zp - None, # block_size - w1_b, - w2_b, - alpha, - limit, - True, # is_vnni - ) - atol = rtol = precision[torch_output.dtype] - torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) + + if bias: + # bias path also exercises clamped SwiGLU (alpha/limit) + w1_b = torch.randn((e, 2 * n), dtype=torch.float32) / 10 + w2_b = torch.randn((e, k), dtype=torch.float32) / 10 + topk_weight, topk_ids = make_routing(m, e, topk, dtype, renormalize) + ref = torch_naive_fused_moe_gptoss( + a, + w1, + w2, + w1_b, + w2_b, + topk_weight, + topk_ids, + False, # already renormalized above when requested + _SWIGLU_ALPHA, + _SWIGLU_LIMIT, + e, + ) + out = run_fused_experts( + a, + packed_w1, + packed_w2, + topk_weight, + topk_ids, + w1_bias=w1_b, + w2_bias=w2_b, + alpha=_SWIGLU_ALPHA, + limit=_SWIGLU_LIMIT, + ) + else: + score = torch.randn((m, e), dtype=dtype) + ref = torch_naive_fused_moe(a, w1, w2, score, topk, renormalize) + topk_weight, topk_ids = kernel.grouped_topk_cpu( + a, score, topk, renormalize, 1, 1, 0, None, None + ) + out = run_fused_experts( + a, packed_w1, packed_w2, topk_weight, topk_ids, inplace=True + ) + + atol = rtol = precision[ref.dtype] + torch.testing.assert_close(ref, out, atol=atol, rtol=rtol) @parametrize(M=[1, 39], N=[128], K=[256], E=[8], topk=[3]) def test_int8_moe(self, M, N, K, E, topk): dtype = torch.bfloat16 - prepack = True - - # Initialize int8 quantization parameters int8_factor_for_scale = 1e-2 int8_max = 127 int8_min = -128 - # Input tensor - # M * K a = torch.randn((M, K), dtype=dtype) / math.sqrt(K) - - # Generate int8 weights w1_fp32 = (torch.rand((E, 2 * N, K), dtype=torch.float32) - 0.5) * 2 w1 = (w1_fp32 * int8_max).clamp(min=int8_min, max=int8_max).to(torch.int8) - w2_fp32 = (torch.rand((E, K, N), dtype=torch.float32) - 0.5) * 2 w2 = (w2_fp32 * int8_max).clamp(min=int8_min, max=int8_max).to(torch.int8) - - # Generate scale for each column (per-column quantization) - w1_s = torch.rand(E, 2 * N, device=w1_fp32.device) * int8_factor_for_scale - w2_s = torch.rand(E, K, device=w2_fp32.device) * int8_factor_for_scale - - # Calculate routing - score = torch.randn((M, E), dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) + w1_s = torch.rand(E, 2 * N) * int8_factor_for_scale + w2_s = torch.rand(E, K) * int8_factor_for_scale + topk_weight, topk_ids = make_routing(M, E, topk, dtype) ref_out = torch_w8a8_per_column_fused_moe( a, w1, w2, w1_s, w2_s, topk_weight, topk_ids, topk ) - - inplace = True - packed_w1 = kernel.convert_weight_packed(w1) if prepack else w1 - packed_w2 = kernel.convert_weight_packed(w2) if prepack else w2 - out = kernel.fused_experts_cpu( + out = run_fused_experts( a, - packed_w1, - packed_w2, + kernel.convert_weight_packed(w1), + kernel.convert_weight_packed(w2), topk_weight, - topk_ids.to(torch.int32), - inplace, - CPUQuantMethod.INT8_W8A8, - w1_s, - w2_s, - None, - None, - None, - None, - None, - None, - None, - prepack, + topk_ids, + quant=CPUQuantMethod.INT8_W8A8, + w1_scale=w1_s, + w2_scale=w2_s, + inplace=True, ) - atol = rtol = precision[ref_out.dtype] - # Increase the tolerance for large input shapes - if M > 35: - atol = rtol = 0.02 + atol = rtol = 0.02 if M > 35 else precision[ref_out.dtype] torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol) - @parametrize(M=[2, 121], N=[352, 512], K=[256, 320], E=[8], topk=[4]) - def test_fp8_moe(self, M, N, K, E, topk): - dtype = torch.bfloat16 - - a = torch.randn(M, K, dtype=dtype) / math.sqrt(K) - - w1_fp32 = torch.randn(E, 2 * N, K) - w1 = (w1_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) - - w2_fp32 = torch.randn(E, K, N) - w2 = (w2_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) - - w1s = ( - torch.randn(E, math.ceil(2 * N / BLOCK_N), math.ceil(K / BLOCK_K)) - * factor_for_scale - ) - w2s = ( - torch.randn(E, math.ceil(K / BLOCK_N), math.ceil(N / BLOCK_K)) - * factor_for_scale - ) - - w1_scaled = scaled_weight(w1, w1s) - w2_scaled = scaled_weight(w2, w2s) - - score = torch.randn((M, E), dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) - - w1 = kernel.convert_weight_packed(w1) - w2 = kernel.convert_weight_packed(w2) - - ref_out = native_fp8_fused_moe( - a, w1_scaled, w2_scaled, topk_weight, topk_ids, topk - ) - out = kernel.fused_experts_cpu( - a, - w1, - w2, - topk_weight, - topk_ids.to(torch.int32), - False, - CPUQuantMethod.FP8_W8A16, - w1s, - w2s, - None, - None, - [BLOCK_N, BLOCK_K], - None, - None, - None, - None, - True, - ) - - atol = rtol = precision[dtype] - torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) - @parametrize( - m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] + M=[2, 32], N=[64, 128], K=[64, 128], E=[4], topk=[2], bias=[False, True] ) - def test_fp8_moe_bias(self, m, n, k, e, topk, renormalize): + def test_fp8_moe(self, M, N, K, E, topk, bias): dtype = torch.bfloat16 - - a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 - - w1_fp32 = torch.randn((e, 2 * n, k), device="cpu", dtype=torch.float32) / 10 - w1 = (w1_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) - w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float32) / 10 - - w2_fp32 = torch.randn((e, k, n), device="cpu", dtype=torch.float32) / 10 - w2 = (w2_fp32 * fp8_max).clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) - w2_b = torch.randn((e, k), device="cpu", dtype=torch.float32) / 10 - + a = torch.randn(M, K, dtype=dtype) / 10 + w1 = (torch.randn(E, 2 * N, K) * fp8_max).clamp(min=fp8_min, max=fp8_max).to( + torch.float8_e4m3fn + ) + w2 = (torch.randn(E, K, N) * fp8_max).clamp(min=fp8_min, max=fp8_max).to( + torch.float8_e4m3fn + ) w1s = ( - torch.randn(e, math.ceil(2 * n / BLOCK_N), math.ceil(k / BLOCK_K)) + torch.randn(E, math.ceil(2 * N / BLOCK_N), math.ceil(K / BLOCK_K)) * factor_for_scale ) w2s = ( - torch.randn(e, math.ceil(k / BLOCK_N), math.ceil(n / BLOCK_K)) + torch.randn(E, math.ceil(K / BLOCK_N), math.ceil(N / BLOCK_K)) * factor_for_scale ) - w1_scaled = scaled_weight(w1, w1s).to(dtype) w2_scaled = scaled_weight(w2, w2s).to(dtype) + topk_weight, topk_ids = make_routing(M, E, topk, dtype) - score = torch.randn((m, e), device="cpu", dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) - alpha = 1.702 - limit = 7.0 - - ref_out = torch_naive_fused_moe_gptoss( - a, - w1_scaled, - w2_scaled, - w1_b, - w2_b, - topk_weight, - topk_ids, - renormalize, - alpha, - limit, - e, + packed_w1 = kernel.convert_weight_packed(w1) + packed_w2 = kernel.convert_weight_packed(w2) + common = dict( + quant=CPUQuantMethod.FP8_W8A16, + w1_scale=w1s, + w2_scale=w2s, + block_size=[BLOCK_N, BLOCK_K], ) - w1 = kernel.convert_weight_packed(w1) - w2 = kernel.convert_weight_packed(w2) - - out = kernel.fused_experts_cpu( - a, - w1, - w2, - topk_weight, - topk_ids.to(torch.int32), - False, - CPUQuantMethod.FP8_W8A16, - w1s, - w2s, - None, - None, - [BLOCK_N, BLOCK_K], - w1_b, - w2_b, - alpha, - limit, - True, - ) + if bias: + w1_b = torch.randn((E, 2 * N), dtype=torch.float32) / 10 + w2_b = torch.randn((E, K), dtype=torch.float32) / 10 + ref = torch_naive_fused_moe_gptoss( + a, + w1_scaled, + w2_scaled, + w1_b, + w2_b, + topk_weight, + topk_ids, + False, + _SWIGLU_ALPHA, + _SWIGLU_LIMIT, + E, + ) + out = run_fused_experts( + a, + packed_w1, + packed_w2, + topk_weight, + topk_ids, + w1_bias=w1_b, + w2_bias=w2_b, + alpha=_SWIGLU_ALPHA, + limit=_SWIGLU_LIMIT, + **common, + ) + else: + ref = native_fp8_fused_moe( + a, w1_scaled.float(), w2_scaled.float(), topk_weight, topk_ids, topk + ) + out = run_fused_experts( + a, packed_w1, packed_w2, topk_weight, topk_ids, **common + ) atol = rtol = precision[dtype] - torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) + torch.testing.assert_close(ref.bfloat16(), out, atol=atol, rtol=rtol) - @parametrize(M=[2, 121], N=[352, 512], K=[256, 320], E=[8], topk=[4]) - def test_mxfp4_moe(self, M, N, K, E, topk): + @parametrize( + M=[2, 32], N=[64, 128], K=[64, 128], E=[4], topk=[2], bias=[False, True] + ) + def test_mxfp4_moe(self, M, N, K, E, topk, bias): dtype = torch.bfloat16 - a = torch.randn(M, K, dtype=dtype) / 10 w1_bf16 = torch.randn((E, 2 * N, K), dtype=dtype) / 10 @@ -341,119 +271,73 @@ def test_mxfp4_moe(self, M, N, K, E, topk): w2s = w2s.reshape(E, K, N // 32) w2dq = MXFP4QuantizeUtil.dequantize(w2q, dtype, w2s) - score = torch.randn((M, E), dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) - - w1 = kernel.convert_weight_packed(w1q) - w2 = kernel.convert_weight_packed(w2q) - w1s = kernel.convert_scale_packed(w1s) - w2s = kernel.convert_scale_packed(w2s) - - ref_out = native_fp8_fused_moe( - a, w1dq.float(), w2dq.float(), topk_weight, topk_ids, topk - ) - out = kernel.fused_experts_cpu( - a, - w1, - w2, - topk_weight, - topk_ids.to(torch.int32), - False, - CPUQuantMethod.MXFP4, - w1s, - w2s, - None, - None, - None, - None, - None, - None, - None, - True, + topk_weight, topk_ids = make_routing(M, E, topk, dtype) + packed_w1 = kernel.convert_weight_packed(w1q) + packed_w2 = kernel.convert_weight_packed(w2q) + packed_w1s = kernel.convert_scale_packed(w1s) + packed_w2s = kernel.convert_scale_packed(w2s) + common = dict( + quant=CPUQuantMethod.MXFP4, + w1_scale=packed_w1s, + w2_scale=packed_w2s, ) - atol = rtol = precision[dtype] - torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) - - @parametrize( - m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] - ) - def test_mxfp4_moe_bias(self, m, n, k, e, topk, renormalize): - dtype = torch.bfloat16 - - a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 - w1_bf16 = torch.randn((e, 2 * n, k), device="cpu", dtype=dtype) / 10 - w1q, w1s = MXFP4QuantizeUtil.quantize(w1_bf16) - w1s = w1s.reshape(e, 2 * n, k // 32) - w1dq = MXFP4QuantizeUtil.dequantize(w1q, dtype, w1s) - w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float32) / 10 - w2_bf16 = torch.randn((e, k, n), device="cpu", dtype=dtype) / 10 - w2q, w2s = MXFP4QuantizeUtil.quantize(w2_bf16) - w2s = w2s.reshape(e, k, n // 32) - w2dq = MXFP4QuantizeUtil.dequantize(w2q, dtype, w2s) - w2_b = torch.randn((e, k), device="cpu", dtype=torch.float32) / 10 - score = torch.randn((m, e), device="cpu", dtype=dtype) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) - alpha = 1.702 - limit = 7.0 - torch_output = torch_naive_fused_moe_gptoss( - a, - w1dq, - w2dq, - w1_b, - w2_b, - topk_weight, - topk_ids, - renormalize, - alpha, - limit, - e, - ) - - w1 = kernel.convert_weight_packed(w1q) - w2 = kernel.convert_weight_packed(w2q) - w1s = kernel.convert_scale_packed(w1s) - w2s = kernel.convert_scale_packed(w2s) - - fused_output = torch.ops.sgl_kernel.fused_experts_cpu( - a, - w1, - w2, - topk_weight, - topk_ids.to(torch.int32), - False, # inplace # See [Note] inplace should be False in fused_experts. - CPUQuantMethod.MXFP4, # use_mxfp4 - w1s, # w1_scale - w2s, # w2_scale - None, # w1_zp - None, # w2_zp - None, # block_size - w1_b, - w2_b, - alpha, - limit, - True, # is_vnni - ) - atol = rtol = precision[torch_output.dtype] - torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) + if bias: + w1_b = torch.randn((E, 2 * N), dtype=torch.float32) / 10 + w2_b = torch.randn((E, K), dtype=torch.float32) / 10 + ref = torch_naive_fused_moe_gptoss( + a, + w1dq, + w2dq, + w1_b, + w2_b, + topk_weight, + topk_ids, + False, + _SWIGLU_ALPHA, + _SWIGLU_LIMIT, + E, + ) + out = run_fused_experts( + a, + packed_w1, + packed_w2, + topk_weight, + topk_ids, + w1_bias=w1_b, + w2_bias=w2_b, + alpha=_SWIGLU_ALPHA, + limit=_SWIGLU_LIMIT, + **common, + ) + torch.testing.assert_close( + ref, out, atol=precision[ref.dtype], rtol=precision[ref.dtype] + ) + else: + ref = native_fp8_fused_moe( + a, w1dq.float(), w2dq.float(), topk_weight, topk_ids, topk + ) + out = run_fused_experts( + a, packed_w1, packed_w2, topk_weight, topk_ids, **common + ) + torch.testing.assert_close( + ref.bfloat16(), out, atol=precision[dtype], rtol=precision[dtype] + ) @parametrize(M=[1, 6], N=[512], K=[256], E=[8], topk=[4]) def test_int4_moe(self, M, N, K, E, topk, group_size=128): dtype = torch.bfloat16 a = torch.rand(M, K, dtype=dtype) / math.sqrt(K) - awq_w13_weight = torch.randint(-127, 128, (E, K, 2 * N // 8)).to(torch.int) awq_w13_zero = torch.randint(0, 10, (E, K // group_size, 2 * N // 8)).to( torch.int ) awq_w13_scales = torch.rand(E, int(K // group_size), 2 * N).to(torch.bfloat16) - awq_w2_weight = torch.randint(-127, 128, (E, N, K // 8)).to(torch.int) awq_w2_zero = torch.randint(0, 10, (E, N // group_size, K // 8)).to(torch.int) awq_w2_scales = torch.rand(E, int(N // group_size), K).to(torch.bfloat16) + bf16_w13_weight = [] bf16_w2_weight = [] for i in range(E): @@ -469,12 +353,10 @@ def test_int4_moe(self, M, N, K, E, topk, group_size=128): bf16_w2_weight = torch.stack(bf16_w2_weight).detach() score = torch.rand((M, E), dtype=dtype) - ref_out = torch_naive_fused_moe( a, bf16_w13_weight, bf16_w2_weight, score, topk, False ) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) + topk_weight, topk_ids = make_routing(M, E, topk, dtype) awq_w13_weight_pack, awq_w13_zero_pack, awq_w13_scales_pack = ( torch.ops.sgl_kernel.convert_weight_packed_scale_zp( awq_w13_weight, awq_w13_zero, awq_w13_scales, 0 @@ -486,24 +368,17 @@ def test_int4_moe(self, M, N, K, E, topk, group_size=128): ) ) - out = kernel.fused_experts_cpu( + out = run_fused_experts( a, awq_w13_weight_pack, awq_w2_weight_pack, topk_weight, - topk_ids.to(torch.int32), - False, - CPUQuantMethod.INT4_W4A8, - awq_w13_scales_pack, - awq_w2_scales_pack, - awq_w13_zero_pack, - awq_w2_zero_pack, - None, - None, - None, - None, - None, - True, + topk_ids, + quant=CPUQuantMethod.INT4_W4A8, + w1_scale=awq_w13_scales_pack, + w2_scale=awq_w2_scales_pack, + w1_zp=awq_w13_zero_pack, + w2_zp=awq_w2_zero_pack, ) atol = rtol = precision[dtype] From 27d7a48590c481ddfba3c6ec9c8b17dfe35353fa Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 16 Jul 2026 08:22:25 +0000 Subject: [PATCH 4/5] Fix test_int4_moe: keep shared routing for ref vs kernel make_routing draws a fresh score, so ref (from the original score) and the kernel no longer selected the same experts. Co-authored-by: Ma Mingfei --- test/registered/cpu/test_moe.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index 39f0e70f2e51..ac1937b90c62 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -352,11 +352,15 @@ def test_int4_moe(self, M, N, K, E, topk, group_size=128): bf16_w13_weight = torch.stack(bf16_w13_weight).detach() bf16_w2_weight = torch.stack(bf16_w2_weight).detach() + # Ref and kernel must share the same routing: torch_naive_fused_moe + # softaxes+topks `score` internally; do not call make_routing here + # (that draws a fresh score and breaks the comparison). score = torch.rand((M, E), dtype=dtype) ref_out = torch_naive_fused_moe( a, bf16_w13_weight, bf16_w2_weight, score, topk, False ) - topk_weight, topk_ids = make_routing(M, E, topk, dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) awq_w13_weight_pack, awq_w13_zero_pack, awq_w13_scales_pack = ( torch.ops.sgl_kernel.convert_weight_packed_scale_zp( awq_w13_weight, awq_w13_zero, awq_w13_scales, 0 From d1109b236f6c4e6928b2224747a4477ed543dfb9 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 16 Jul 2026 08:34:26 +0000 Subject: [PATCH 5/5] Fix FP8 MoE bias test: keep float32 scaled weights for ref Casting scaled_weight() to bf16 truncates relative to the FP8 kernel's float accumulation and breaks test_fp8_moe_bias at larger shapes (e.g. m=32, n=k=128). Match main's test_fp8_moe and keep float32 dequant weights in the reference path. Co-authored-by: Ma Mingfei --- test/registered/cpu/test_moe.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index ac1937b90c62..5daf592e3a41 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -202,8 +202,11 @@ def test_fp8_moe(self, M, N, K, E, topk, bias): torch.randn(E, math.ceil(K / BLOCK_N), math.ceil(N / BLOCK_K)) * factor_for_scale ) - w1_scaled = scaled_weight(w1, w1s).to(dtype) - w2_scaled = scaled_weight(w2, w2s).to(dtype) + # Keep float32 dequant weights for the reference. Casting to bf16 here + # truncates vs the FP8 kernel's float accumulate and makes bias/swiglu + # cases (e.g. m=32, n=k=128) exceed precision[bf16]. + w1_scaled = scaled_weight(w1, w1s) + w2_scaled = scaled_weight(w2, w2s) topk_weight, topk_ids = make_routing(M, E, topk, dtype) packed_w1 = kernel.convert_weight_packed(w1) @@ -245,7 +248,7 @@ def test_fp8_moe(self, M, N, K, E, topk, bias): ) else: ref = native_fp8_fused_moe( - a, w1_scaled.float(), w2_scaled.float(), topk_weight, topk_ids, topk + a, w1_scaled, w2_scaled, topk_weight, topk_ids, topk ) out = run_fused_experts( a, packed_w1, packed_w2, topk_weight, topk_ids, **common