From a02b337e0ff998865d9238fb3f4d36d7b11ac798 Mon Sep 17 00:00:00 2001 From: gc-fu Date: Mon, 13 Jul 2026 02:47:00 +0000 Subject: [PATCH] fix(esimd): declare INT4 output mutations --- .../csrc/xpu/torch_extension.cc | 44 ++++++++++++++----- .../python/custom_esimd_kernels_vllm/ops.py | 10 ++--- 2 files changed, 38 insertions(+), 16 deletions(-) diff --git a/vllm/custom-esimd-kernels-vllm/csrc/xpu/torch_extension.cc b/vllm/custom-esimd-kernels-vllm/csrc/xpu/torch_extension.cc index 326e65eb..0e2fa1e3 100644 --- a/vllm/custom-esimd-kernels-vllm/csrc/xpu/torch_extension.cc +++ b/vllm/custom-esimd-kernels-vllm/csrc/xpu/torch_extension.cc @@ -46,14 +46,22 @@ TORCH_LIBRARY(custom_esimd_kernels_vllm, m) { // INT4 GEMV with per-group scale (group_size=128) // Weight [N, K/2] uint8 packed, scale [N, K/128] fp16. N/K auto-detected. m.def("esimd_gemv_int4(Tensor input, Tensor weight, Tensor weight_scale, " - "Tensor output) -> Tensor"); - m.impl("esimd_gemv_int4", torch::kXPU, &esimd_gemv_int4); + "Tensor(a!) output) -> ()"); + m.impl("esimd_gemv_int4", torch::kXPU, + [](at::Tensor input, at::Tensor weight, at::Tensor weight_scale, + at::Tensor output) -> void { + esimd_gemv_int4(input, weight, weight_scale, output); + }); // Fused 2-matrix INT4 GEMV (GDN in_proj_qkvz + in_proj_ba) m.def("esimd_gemv_int4_fused2(Tensor input, " - "Tensor w0, Tensor s0, Tensor o0, " - "Tensor w1, Tensor s1, Tensor o1) -> Tensor"); - m.impl("esimd_gemv_int4_fused2", torch::kXPU, &esimd_gemv_int4_fused2); + "Tensor w0, Tensor s0, Tensor(a!) o0, " + "Tensor w1, Tensor s1, Tensor(b!) o1) -> ()"); + m.impl("esimd_gemv_int4_fused2", torch::kXPU, + [](at::Tensor input, at::Tensor w0, at::Tensor s0, at::Tensor o0, + at::Tensor w1, at::Tensor s1, at::Tensor o1) -> void { + esimd_gemv_int4_fused2(input, w0, s0, o0, w1, s1, o1); + }); // Fused QKV Split + RMSNorm + RoPE // q_out/gate_out/k_out/v_out are written in place (mutated). Declared with @@ -116,16 +124,30 @@ TORCH_LIBRARY(custom_esimd_kernels_vllm, m) { m.impl("esimd_norm_gemv_fp8_pert", torch::kXPU, &esimd_norm_gemv_fp8_pert); // Fused ResidualAdd + RMSNorm + INT4 GEMV (post_attn_norm + router) - m.def("esimd_resadd_norm_gemv_int4_pert(Tensor hidden_states, Tensor residual, " + m.def("esimd_resadd_norm_gemv_int4_pert(Tensor hidden_states, Tensor(a!) residual, " "Tensor norm_weight, Tensor gemv_weight, Tensor gemv_scale, " - "Tensor output, Tensor normed_out, float eps) -> Tensor"); - m.impl("esimd_resadd_norm_gemv_int4_pert", torch::kXPU, &esimd_resadd_norm_gemv_int4_pert); + "Tensor(b!) output, Tensor(c!) normed_out, float eps) -> ()"); + m.impl("esimd_resadd_norm_gemv_int4_pert", torch::kXPU, + [](at::Tensor hidden_states, at::Tensor residual, + at::Tensor norm_weight, at::Tensor gemv_weight, + at::Tensor gemv_scale, at::Tensor output, at::Tensor normed_out, + double eps) -> void { + esimd_resadd_norm_gemv_int4_pert(hidden_states, residual, norm_weight, + gemv_weight, gemv_scale, output, + normed_out, eps); + }); // Fused RMSNormGated + INT4 GEMV (out_proj for GDN layers) m.def("esimd_norm_gemv_int4_pert(Tensor x, Tensor z, Tensor norm_weight, " - "Tensor gemv_weight, Tensor gemv_scale, Tensor output, " - "int HV, int V, float eps) -> Tensor"); - m.impl("esimd_norm_gemv_int4_pert", torch::kXPU, &esimd_norm_gemv_int4_pert); + "Tensor gemv_weight, Tensor gemv_scale, Tensor(a!) output, " + "int HV, int V, float eps) -> ()"); + m.impl("esimd_norm_gemv_int4_pert", torch::kXPU, + [](at::Tensor x, at::Tensor z, at::Tensor norm_weight, + at::Tensor gemv_weight, at::Tensor gemv_scale, at::Tensor output, + int64_t HV, int64_t V, double eps) -> void { + esimd_norm_gemv_int4_pert(x, z, norm_weight, gemv_weight, + gemv_scale, output, HV, V, eps); + }); // Standalone RMSNorm (no add): for the spots where input differs from // the accumulating residual stream (post_attn_norm, post_ff_norm_1, etc.). diff --git a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py index 60a759c5..0b24720b 100644 --- a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py +++ b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py @@ -108,7 +108,7 @@ def esimd_gemv_fp8_pert_fused3( def esimd_gemv_int4( input: torch.Tensor, weight: torch.Tensor, weight_scale: torch.Tensor, output: torch.Tensor, -) -> torch.Tensor: +) -> None: """Symmetric INT4 weight GEMV with per-group scale, FP32 accumulation. Computes: output[1, N] = input[1, K] @ dequant(weight)^T @@ -130,7 +130,7 @@ def esimd_gemv_int4_fused2( input: torch.Tensor, w0: torch.Tensor, s0: torch.Tensor, o0: torch.Tensor, w1: torch.Tensor, s1: torch.Tensor, o1: torch.Tensor, -) -> torch.Tensor: +) -> None: """Fused 2-matrix INT4 GEMV: two GEMVs sharing the same input, single kernel. Saves one kernel launch overhead (~20-50 us) compared to two separate calls. @@ -140,7 +140,7 @@ def esimd_gemv_int4_fused2( w0: [N0, K/2] uint8, s0: [N0, K/128] fp16, o0: [1, N0] fp16 w1: [N1, K/2] uint8, s1: [N1, K/128] fp16, o1: [1, N1] fp16 - Returns o0. Both o0 and o1 are written. + Both o0 and o1 are written in place. """ return _ops.esimd_gemv_int4_fused2(input, w0, s0, o0, w1, s1, o1) @@ -410,7 +410,7 @@ def esimd_resadd_norm_gemv_int4_pert( output: torch.Tensor, normed_out: torch.Tensor, eps: float, -) -> torch.Tensor: +) -> None: """Fused ResidualAdd + RMSNorm + INT4 GEMV. Combines post_attention_layernorm + MoE router GEMV (INT4 quantized): @@ -488,7 +488,7 @@ def esimd_norm_gemv_int4_pert( HV: int, V: int, eps: float, -) -> torch.Tensor: +) -> None: """Fused RMSNormGated + INT4 GEMV for GDN out_proj decode path. Combines norm(x, z) + out_proj(normed) into a single kernel.