[CPU] Add graph register for fused_sigmoid_mul_cpu, fused_qk_gemma_rmsnorm (#35506)
This commit is contained in:
@@ -173,7 +173,7 @@ at::Tensor gelu_and_mul_cpu(const at::Tensor& input) {
|
||||
return out;
|
||||
}
|
||||
|
||||
at::Tensor fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate, bool inplace) {
|
||||
void fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate) {
|
||||
CHECK_DIM(2, input);
|
||||
const int64_t gate_dim = gate.dim();
|
||||
TORCH_CHECK(gate_dim == 2 || gate_dim == 3, "gate must be a 2D or 3D tensor");
|
||||
@@ -195,10 +195,9 @@ at::Tensor fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate, bool
|
||||
int64_t g_strideT = gate.stride(0);
|
||||
int64_t g_strideH = is_gate_3d ? gate.stride(1) : 0;
|
||||
|
||||
at::Tensor out = inplace ? input : at::empty_like(input);
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "fused_sigmoid_mul", [&] {
|
||||
fused_sigmoid_mul_kernel_impl<scalar_t>(
|
||||
out.data_ptr<scalar_t>(),
|
||||
input.data_ptr<scalar_t>(),
|
||||
input.data_ptr<scalar_t>(),
|
||||
gate.data_ptr<scalar_t>(),
|
||||
num_tokens,
|
||||
@@ -208,5 +207,4 @@ at::Tensor fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate, bool
|
||||
g_strideT,
|
||||
g_strideH);
|
||||
});
|
||||
return out;
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ at::Tensor gelu_tanh_and_mul_cpu(const at::Tensor& input);
|
||||
at::Tensor gelu_and_mul_cpu(const at::Tensor& input);
|
||||
|
||||
// fused_sigmoid_mul
|
||||
at::Tensor fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate, bool inplace);
|
||||
void fused_sigmoid_mul_cpu(at::Tensor& input, const at::Tensor& gate);
|
||||
|
||||
// l2norm
|
||||
at::Tensor l2norm_cpu(at::Tensor& input, double eps);
|
||||
@@ -579,7 +579,7 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
m.impl("gelu_tanh_and_mul_cpu", torch::kCPU, &gelu_tanh_and_mul_cpu);
|
||||
m.def("gelu_and_mul_cpu(Tensor input) -> Tensor");
|
||||
m.impl("gelu_and_mul_cpu", torch::kCPU, &gelu_and_mul_cpu);
|
||||
m.def("fused_sigmoid_mul_cpu(Tensor(a!) input, Tensor gate, bool inplace) -> Tensor(a!)");
|
||||
m.def("fused_sigmoid_mul_cpu(Tensor(a!) input, Tensor gate) -> ()");
|
||||
m.impl("fused_sigmoid_mul_cpu", torch::kCPU, &fused_sigmoid_mul_cpu);
|
||||
|
||||
// norm
|
||||
|
||||
@@ -185,6 +185,7 @@ def register_fake_ops(tp_size: int):
|
||||
"fused_add_layernorm_cpu",
|
||||
"multimodal_rotary_embedding_cpu",
|
||||
"apply_multidimensional_rope_cpu",
|
||||
"fused_sigmoid_mul_cpu",
|
||||
]
|
||||
for op in none_return_ops:
|
||||
|
||||
@@ -209,6 +210,19 @@ def register_fake_ops(tp_size: int):
|
||||
def _(input, *args, **kwargs):
|
||||
return torch.empty_like(input)
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_gemma_rmsnorm_cpu")
|
||||
def _(q, k, q_weight, k_weight, eps, head_dim):
|
||||
return torch.empty_like(q), torch.empty_like(k)
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_gemma_rmsnorm_with_gate_cpu")
|
||||
def _(q_gate, k, q_weight, k_weight, eps, head_dim, num_head):
|
||||
seq_len = q_gate.shape[0]
|
||||
num_head_kv = k.shape[1] // head_dim
|
||||
q_out = q_gate.new_empty((seq_len * num_head, head_dim))
|
||||
k_out = k.new_empty((seq_len * num_head_kv, head_dim))
|
||||
gate_out = q_gate.new_empty((seq_len * num_head, head_dim))
|
||||
return q_out, k_out, gate_out
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_rmsnorm_cpu")
|
||||
def _(q, k, *args, **kwargs):
|
||||
return torch.empty_like(q), torch.empty_like(k)
|
||||
|
||||
@@ -190,7 +190,14 @@ if _is_cuda:
|
||||
)
|
||||
|
||||
if _is_cpu:
|
||||
fused_sigmoid_mul = torch.ops.sgl_kernel.fused_sigmoid_mul_cpu
|
||||
_fused_sigmoid_mul_cpu = torch.ops.sgl_kernel.fused_sigmoid_mul_cpu
|
||||
|
||||
def fused_sigmoid_mul(x, gate, inplace=True):
|
||||
if not inplace:
|
||||
x = x.clone()
|
||||
_fused_sigmoid_mul_cpu(x, gate)
|
||||
return x
|
||||
|
||||
fused_qk_gemma_rmsnorm = torch.ops.sgl_kernel.fused_qk_gemma_rmsnorm_cpu
|
||||
fused_qk_gemma_rmsnorm_with_gate = (
|
||||
torch.ops.sgl_kernel.fused_qk_gemma_rmsnorm_with_gate_cpu
|
||||
|
||||
Reference in New Issue
Block a user