diff --git a/python/sglang/kernels/aot/csrc/cpu/activation.cpp b/python/sglang/kernels/aot/csrc/cpu/activation.cpp index 658577bc1..29355e74f 100644 --- a/python/sglang/kernels/aot/csrc/cpu/activation.cpp +++ b/python/sglang/kernels/aot/csrc/cpu/activation.cpp @@ -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( - out.data_ptr(), + input.data_ptr(), input.data_ptr(), gate.data_ptr(), 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; } diff --git a/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp b/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp index e9aca568b..573674ab2 100644 --- a/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp +++ b/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp @@ -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 diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 1367ee17e..102828d7a 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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) diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index 8d0c9e429..3371ac0d3 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -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 diff --git a/test/registered/cpu/test_activation.py b/test/registered/cpu/test_activation.py index 4ca87f45e..d8e966494 100644 --- a/test/registered/cpu/test_activation.py +++ b/test/registered/cpu/test_activation.py @@ -50,6 +50,8 @@ def test_activation(m, n, dtype): @pytest.mark.parametrize("num_heads", [16]) @pytest.mark.parametrize("m", [1, 17, 128]) def test_fused_sigmoid_mul(m, num_heads, head_dim, dtype, gate_3d): + from sglang.srt.models.qwen3_5 import fused_sigmoid_mul + x = torch.randn([m, num_heads * head_dim], dtype=dtype) if gate_3d: gate_storage = torch.randn([m, num_heads, head_dim * 2], dtype=dtype) @@ -59,16 +61,23 @@ def test_fused_sigmoid_mul(m, num_heads, head_dim, dtype, gate_3d): gate = torch.randn_like(x) gate_ref = gate.reshape(m, -1) if gate_3d else gate - _assert_close( - x * torch.sigmoid(gate_ref), - torch.ops.sgl_kernel.fused_sigmoid_mul_cpu(x, gate, False), - ) x_inplace = x.clone() - ref_inplace = x_inplace * torch.sigmoid(gate_ref) - out_inplace = torch.ops.sgl_kernel.fused_sigmoid_mul_cpu(x_inplace, gate, True) - assert out_inplace.data_ptr() == x_inplace.data_ptr() - _assert_close(ref_inplace, x_inplace) + ref = x_inplace * torch.sigmoid(gate_ref) + torch.ops.sgl_kernel.fused_sigmoid_mul_cpu(x_inplace, gate) + _assert_close(ref, x_inplace) + + x_out_of_place = x.clone() + out = fused_sigmoid_mul(x_out_of_place, gate, inplace=False) + assert out.data_ptr() != x_out_of_place.data_ptr() + assert out.data_ptr() != x.data_ptr() + _assert_close(ref, out) + _assert_close(x, x_out_of_place) + + x_wrapper_inplace = x.clone() + out_inplace = fused_sigmoid_mul(x_wrapper_inplace, gate, inplace=True) + assert out_inplace.data_ptr() == x_wrapper_inplace.data_ptr() + _assert_close(ref, x_wrapper_inplace) if __name__ == "__main__":