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..275a06b39e50 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) { @@ -217,63 +228,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_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/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 diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index 0ede6aa40c0a..5daf592e3a41 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -32,180 +32,168 @@ 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): + @parametrize( + M=[2, 32], N=[64, 128], K=[64, 128], E=[4], topk=[2], bias=[False, True] + ) + def test_fp8_moe(self, M, N, K, E, topk, bias): 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) - + 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)) * factor_for_scale @@ -214,47 +202,66 @@ def test_fp8_moe(self, M, N, K, E, topk): torch.randn(E, math.ceil(K / BLOCK_N), math.ceil(N / BLOCK_K)) * factor_for_scale ) - + # 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) - 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, + 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], ) + 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, w2_scaled, 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 @@ -267,119 +274,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, - ) - - 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, + 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, ) - 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): @@ -394,8 +355,10 @@ 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 ) @@ -412,24 +375,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]