Add Arm64 CPU Phase 1A CI bootstrap (#22123)

Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
Mandepudi Rani Chowdary
2026-05-08 09:28:23 +08:00
committed by GitHub
co-authored by Ma Mingfei
parent 3c3f0bd55e
commit 55224fff08
7 changed files with 272 additions and 22 deletions
+17
View File
@@ -75,6 +75,23 @@ endif()
file(GLOB_RECURSE SOURCES "${CMAKE_CURRENT_SOURCE_DIR}/*.cpp")
# These kernels still rely on x86-specific AMX/AVX512 implementations.
# Keep them out of Arm64 bootstrap builds until native Arm paths land.
set(SGLANG_CPU_X86_ONLY_SOURCES
${CMAKE_CURRENT_SOURCE_DIR}/gemm_int4.cpp
${CMAKE_CURRENT_SOURCE_DIR}/moe.cpp
${CMAKE_CURRENT_SOURCE_DIR}/moe_fp8.cpp
${CMAKE_CURRENT_SOURCE_DIR}/moe_int4.cpp
${CMAKE_CURRENT_SOURCE_DIR}/moe_int8.cpp
${CMAKE_CURRENT_SOURCE_DIR}/qkv_proj.cpp
${CMAKE_CURRENT_SOURCE_DIR}/mamba/conv.cpp
)
if(CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64|arm64")
add_compile_definitions(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS)
list(REMOVE_ITEM SOURCES ${SGLANG_CPU_X86_ONLY_SOURCES})
endif()
if(NOT DEFINED ENV{SGLANG_CPU_FP8_CVT_FTZ})
set(ENV{SGLANG_CPU_FP8_CVT_FTZ} "1")
endif()
+29 -21
View File
@@ -150,18 +150,6 @@ at::Tensor convert_scale_packed(at::Tensor& scale);
// quant
std::tuple<at::Tensor, at::Tensor> per_token_quant_int8_cpu(at::Tensor& A);
// gemm
at::Tensor
weight_packed_linear(at::Tensor& mat1, at::Tensor& mat2, const std::optional<at::Tensor>& bias, bool is_vnni);
// gemm fusion
at::Tensor fused_linear_sigmoid_mul(
at::Tensor& mat1,
at::Tensor& mat2,
const std::optional<at::Tensor>& bias,
bool is_vnni,
const at::Tensor& post_mul_mat);
// igemm
at::Tensor int8_scaled_mm_cpu(
at::Tensor& mat1,
@@ -195,6 +183,7 @@ at::Tensor int8_scaled_mm_with_quant(
at::ScalarType out_dtype,
bool is_vnni);
#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS)
// int4 gemm
at::Tensor int4_scaled_mm_cpu(
at::Tensor& x, at::Tensor& w, at::Tensor& w_zeros, at::Tensor& w_scales, std::optional<at::Tensor> bias);
@@ -205,10 +194,24 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> convert_weight_packed_scale_zp(
at::Tensor qzeros, // awq: (*, K / group_size, N / 8) || gptq: (*, K / group_size, N / 8) , int32
at::Tensor scales, // awq: (*, K / group_size, N) || gptq: (*, K / group_size, N) , bfloat16
int64_t quant_method_4bit);
#endif
// gemm
at::Tensor
weight_packed_linear(at::Tensor& mat1, at::Tensor& mat2, const std::optional<at::Tensor>& bias, bool is_vnni);
// gemm fusion
at::Tensor fused_linear_sigmoid_mul(
at::Tensor& mat1,
at::Tensor& mat2,
const std::optional<at::Tensor>& bias,
bool is_vnni,
const at::Tensor& post_mul_mat);
// bmm
void bmm_cpu(at::Tensor& out, at::Tensor& mat1, at::Tensor& mat2, bool is_vnni, const std::optional<at::Tensor>& scale);
#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS)
// fused moe
at::Tensor fused_experts_cpu(
at::Tensor& hidden_states,
@@ -306,6 +309,7 @@ at::Tensor causal_conv1d_update_cpu(
const std::optional<at::Tensor>& conv_state_indices,
int64_t pad_slot_id,
bool is_vnni);
#endif
// conv3d fast path for patch embedding
at::Tensor conv3d_embed_weight_pack(const at::Tensor& weight);
@@ -495,15 +499,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
m.def("per_token_quant_int8_cpu(Tensor A) -> (Tensor, Tensor)");
m.impl("per_token_quant_int8_cpu", torch::kCPU, &per_token_quant_int8_cpu);
// gemm
m.def("weight_packed_linear(Tensor mat1, Tensor mat2, Tensor? bias, bool is_vnni) -> Tensor");
m.impl("weight_packed_linear", torch::kCPU, &weight_packed_linear);
// gemm fusion
m.def(
"fused_linear_sigmoid_mul(Tensor mat1, Tensor mat2, Tensor? bias, bool is_vnni, Tensor post_mul_mat) -> Tensor");
m.impl("fused_linear_sigmoid_mul", torch::kCPU, &fused_linear_sigmoid_mul);
// igemm
m.def(
"int8_scaled_mm_cpu(Tensor mat1, Tensor mat2, Tensor scales1, Tensor scales2, Tensor? bias, ScalarType "
@@ -526,6 +521,7 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
"is_vnni) -> Tensor");
m.impl("int8_scaled_mm_with_quant", torch::kCPU, &int8_scaled_mm_with_quant);
#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS)
// int4 gemm
m.def("int4_scaled_mm_cpu(Tensor x, Tensor w, Tensor w_zeros, Tensor w_scales, Tensor? bias) -> Tensor");
m.impl("int4_scaled_mm_cpu", torch::kCPU, &int4_scaled_mm_cpu);
@@ -535,11 +531,22 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
"convert_weight_packed_scale_zp(Tensor weight, Tensor qzeros, Tensor scales, int quant_method_4bit) -> (Tensor, "
"Tensor, Tensor)");
m.impl("convert_weight_packed_scale_zp", torch::kCPU, &convert_weight_packed_scale_zp);
#endif
// gemm
m.def("weight_packed_linear(Tensor mat1, Tensor mat2, Tensor? bias, bool is_vnni) -> Tensor");
m.impl("weight_packed_linear", torch::kCPU, &weight_packed_linear);
// gemm fusion
m.def(
"fused_linear_sigmoid_mul(Tensor mat1, Tensor mat2, Tensor? bias, bool is_vnni, Tensor post_mul_mat) -> Tensor");
m.impl("fused_linear_sigmoid_mul", torch::kCPU, &fused_linear_sigmoid_mul);
// bmm
m.def("bmm_cpu(Tensor(a!) out, Tensor mat1, Tensor mat2, bool is_vnni, Tensor? scale) -> ()");
m.impl("bmm_cpu", torch::kCPU, &bmm_cpu);
#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS)
// moe
m.def(
"fused_experts_cpu(Tensor hidden_states, Tensor w1, Tensor w2, Tensor topk_weights, Tensor topk_ids, bool "
@@ -585,6 +592,7 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
"causal_conv1d_update_cpu(Tensor x, Tensor(a!) conv_states, Tensor weight, Tensor? bias, bool silu_activation,"
"Tensor? cache_seqlens, Tensor? conv_state_indices, int pad_slot_id, bool is_vnni) -> Tensor");
m.impl("causal_conv1d_update_cpu", torch::kCPU, &causal_conv1d_update_cpu);
#endif
// conv3d fast path for patch embedding
m.def("conv3d_embed_weight_pack(Tensor weight) -> Tensor");