[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
@@ -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)");