[CPU] add fused_qk_gemma_norm and refactor norm kernel implementation (#30216)

This commit is contained in:
Ma Mingfei
2026-07-07 08:52:59 +08:00
committed by GitHub
parent 6c1fb8a937
commit 30fb0dd851
6 changed files with 983 additions and 1131 deletions
File diff suppressed because it is too large Load Diff
@@ -58,6 +58,23 @@ at::Tensor fused_add_layernorm_cpu(
const std::optional<at::Tensor>& bias,
double eps);
// fused_qk_gemma_rmsnorm
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
const at::Tensor& q,
const at::Tensor& k,
const at::Tensor& q_weight,
const at::Tensor& k_weight,
double eps,
int64_t head_dim);
std::tuple<at::Tensor, at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_with_gate_cpu(
const at::Tensor& q_gate,
const at::Tensor& k,
const at::Tensor& q_weight,
const at::Tensor& k_weight,
double eps,
int64_t head_dim,
int64_t num_head);
// topk
std::tuple<at::Tensor, at::Tensor>
topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize);
@@ -468,6 +485,15 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
"fused_add_layernorm_cpu(Tensor input, Tensor residual, Tensor weight, Tensor? bias, float eps) -> "
"Tensor");
m.impl("fused_add_layernorm_cpu", torch::kCPU, &fused_add_layernorm_cpu);
m.def(
"fused_qk_gemma_rmsnorm_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, float eps, int head_dim) -> "
"(Tensor, Tensor)");
m.impl("fused_qk_gemma_rmsnorm_cpu", torch::kCPU, &fused_qk_gemma_rmsnorm_cpu);
m.def(
"fused_qk_gemma_rmsnorm_with_gate_cpu(Tensor q_gate, Tensor k, Tensor q_weight, Tensor k_weight, float eps, int "
"head_dim, int num_head) -> "
"(Tensor, Tensor, Tensor)");
m.impl("fused_qk_gemma_rmsnorm_with_gate_cpu", torch::kCPU, &fused_qk_gemma_rmsnorm_with_gate_cpu);
// topk
m.def("topk_sigmoid_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize) -> (Tensor, Tensor)");
+56
View File
@@ -451,6 +451,62 @@ inline std::tuple<__m512i, __m512i> transpose_2x32_16bit(__m512i r0, __m512i r1)
}
#pragma GCC diagnostic pop
// Note: mapped from aten exp_u20
inline __attribute__((always_inline)) __m512 _mm512_exp_u20_ps(const __m512 values) {
const __m512 vec_factorial_1 = _mm512_set1_ps(0.999999701f);
const __m512 vec_factorial_2 = _mm512_set1_ps(0.499991506f);
const __m512 vec_factorial_3 = _mm512_set1_ps(0.166676521f);
const __m512 vec_factorial_4 = _mm512_set1_ps(0.0418978221f);
const __m512 vec_factorial_5 = _mm512_set1_ps(0.00828929059f);
const __m512 vec_exp_log2ef = _mm512_castsi512_ps(_mm512_set1_epi32(0x3fb8aa3b)); // log2(e)
const __m512 vec_half = _mm512_set1_ps(0.5f);
const __m512 vec_one = _mm512_set1_ps(1.f);
const __m512 vec_zero = _mm512_set1_ps(0.f);
const __m512 vec_two = _mm512_set1_ps(2.f);
const __m512 vec_ln2f = _mm512_castsi512_ps(_mm512_set1_epi32(0x3f317218));
const __m512 vec_ln_flt_min = _mm512_castsi512_ps(_mm512_set1_epi32(0xc2aeac50));
const __m512 vec_ln_flt_max = _mm512_castsi512_ps(_mm512_set1_epi32(0x42b17218));
const __m512i vec_127 = _mm512_set1_epi32(0x0000007f);
const int n_mantissa_bits = 23;
// exp(x) =
// = exp(n * ln(2) + r) // divide x by ln(2) and get quot and rem
// = 2^n * exp(r) // simplify the exp(n*ln(2)) expression
auto less_ln_flt_min_mask = _mm512_cmp_ps_mask(values, vec_ln_flt_min, 1 /*_CMP_LT_OS*/);
auto vec_src = _mm512_min_ps(values, vec_ln_flt_max);
vec_src = _mm512_max_ps(vec_src, vec_ln_flt_min);
// fx = floorf(x * log2ef + 0.5)
auto vec_fx = _mm512_fmadd_ps(vec_src, vec_exp_log2ef, vec_half);
auto vec_fx_i = _mm512_cvt_roundps_epi32(vec_fx, _MM_FROUND_TO_NEG_INF | _MM_FROUND_NO_EXC);
vec_fx = _mm512_cvtepi32_ps(vec_fx_i);
// x = x - fx * ln2
auto vec_exp_poly = _mm512_fnmadd_ps(vec_fx, vec_ln2f, vec_src);
// compute polynomial
auto vec_res = _mm512_fmadd_ps(vec_exp_poly, vec_factorial_5, vec_factorial_4);
vec_res = _mm512_fmadd_ps(vec_exp_poly, vec_res, vec_factorial_3);
vec_res = _mm512_fmadd_ps(vec_exp_poly, vec_res, vec_factorial_2);
vec_res = _mm512_fmadd_ps(vec_exp_poly, vec_res, vec_factorial_1);
vec_res = _mm512_fmadd_ps(vec_exp_poly, vec_res, vec_one);
// compute 2^(n-1)
auto vec_exp_number = _mm512_sub_ps(vec_fx, vec_one);
auto vec_exp_number_i = _mm512_cvtps_epi32(vec_exp_number);
auto vec_two_pow_n_i = _mm512_add_epi32(vec_exp_number_i, vec_127);
vec_two_pow_n_i = _mm512_slli_epi32(vec_two_pow_n_i, n_mantissa_bits);
auto vec_two_pow_n = _mm512_castsi512_ps(vec_two_pow_n_i);
vec_two_pow_n = _mm512_mask_blend_ps(less_ln_flt_min_mask, vec_two_pow_n, vec_zero);
// y = y * 2^n
vec_res = _mm512_mul_ps(vec_res, vec_two_pow_n);
vec_res = _mm512_mul_ps(vec_res, vec_two);
return vec_res;
}
// Note: mapped from aten fexp_u20
inline __attribute__((always_inline)) __m512 _mm512_fexp_u20_ps(const __m512 values) {
const __m512 vec_c0 = _mm512_set1_ps(0.00010703434948458272f);
const __m512 vec_c1 = _mm512_set1_ps(0.30354260500649682f);