[CPU] Add missing scalar fallback for CPU W4A8 INT4 GEMM (#44523)

Signed-off-by: wcy <233313160abc@gmail.com>
Co-authored-by: lyd1992 <liuyudong@iscas.ac.cn>
This commit is contained in:
wcy
2026-06-11 08:52:01 +00:00
committed by GitHub
co-authored by lyd1992
parent 1c3a72b8b2
commit f06aefb4e3
4 changed files with 61 additions and 15 deletions
+6
View File
@@ -438,6 +438,12 @@ if(USE_ONEDNN)
${VLLM_EXT_SRC})
endif()
if (CMAKE_SYSTEM_PROCESSOR MATCHES "riscv64")
set(VLLM_EXT_SRC
"csrc/cpu/sgl-kernels/gemm_int4.cpp"
${VLLM_EXT_SRC})
endif()
if (ENABLE_X86_ISA)
set(VLLM_EXT_SRC_SGL
"csrc/cpu/sgl-kernels/conv.cpp"
+37 -1
View File
@@ -268,6 +268,23 @@ void _dequant_gemm_accum_small_M(
_dequant_gemm_accum_small_M<M, N, ldb, sym_quant_act>(C, A, scales_a, qzeros_a, B, scales_b, qzeros_b, K, lda, ldc);
#endif
template <int64_t N, int64_t ldb>
inline int32_t load_uint4_vnni(const uint8_t* __restrict__ B, int64_t k, int64_t n) {
// B is packed as [_block_k / 4, N / 2, 4] for VNNI4. Each byte stores two
// columns from adjacent 8-column groups for one K lane.
constexpr int64_t n_group_size = 8;
constexpr int64_t vnni_size = 4;
static_assert(N % (2 * n_group_size) == 0);
int64_t n_group = n / n_group_size;
int64_t ni = n % n_group_size;
int64_t ki = k % vnni_size;
int64_t k_base = k - ki;
int64_t packed_n = (n_group / 2) * n_group_size + ni;
uint8_t packed = B[k_base * ldb + packed_n * vnni_size + ki];
return (n_group % 2 == 0) ? (packed & 0x0f) : ((packed >> 4) & 0x0f);
}
template <int64_t N, int64_t ldb, bool sym_quant_act>
void _dequant_gemm_accum(
float* C,
@@ -321,7 +338,24 @@ void _dequant_gemm_accum(
} else
#endif
{
TORCH_CHECK(false, "tinygemm_kernel: scalar path not implemented!");
for (int64_t m = 0; m < M; ++m) {
for (int64_t n = 0; n < N; ++n) {
int32_t acc = 0;
for (int64_t k = 0; k < K; ++k) {
int32_t b = load_uint4_vnni<N, ldb>(B, k, n) - qzeros_b[n];
if constexpr (sym_quant_act) {
const int8_t* A_s8 = reinterpret_cast<const int8_t*>(A);
acc += static_cast<int32_t>(A_s8[m * lda + k]) * b;
} else {
acc += static_cast<int32_t>(A[m * lda + k]) * b;
}
}
if constexpr (!sym_quant_act) {
acc -= qzeros_a[m] * compensation[n];
}
C[m * ldc + n] += static_cast<float>(acc) * scales_a[m] * scales_b[n];
}
}
}
}
@@ -496,9 +530,11 @@ void _da8w4_linear_impl(
store_out<out_dtype, BLOCK_N>(C_tmp, output + mci * block_m * N + nc * BLOCK_N, m_size, N /*lda*/);
}
}
#if defined(CPU_CAPABILITY_AVX512)
if (use_brgemm) {
at::native::cpublas::brgemm_release();
}
#endif
});
}
+1 -1
View File
@@ -245,7 +245,7 @@ quantize_row_int8(uint8_t* __restrict__ Aq, float& As, const scalar_t* __restric
for (int64_t k = 0; k < K; ++k) {
const float val = static_cast<float>(A[k]) * inv_scale;
Aq[k] = (uint8_t)(std::round(val)) + 128;
Aq[k] = static_cast<uint8_t>(static_cast<int32_t>(std::round(val)) + 128);
}
As = scale;
}
+17 -13
View File
@@ -429,19 +429,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
ops.impl("int8_scaled_mm_with_quant", torch::kCPU,
&int8_scaled_mm_with_quant);
// Adapted from sglang: INT4 W4A8 kernels
ops.def(
"convert_weight_packed_scale_zp(Tensor weight, Tensor qzeros, Tensor "
"scales, int quant_method_4bit) -> (Tensor, "
"Tensor, Tensor)");
ops.impl("convert_weight_packed_scale_zp", torch::kCPU,
&convert_weight_packed_scale_zp);
ops.def(
"int4_scaled_mm_cpu(Tensor(a0!) x, Tensor(a1!) w, Tensor(a2!) w_zeros, "
"Tensor(a3!) w_scales, Tensor? bias) -> Tensor");
ops.impl("int4_scaled_mm_cpu", torch::kCPU, &int4_scaled_mm_cpu);
// Adapted from sglang: FP8 W8A16 kernel
ops.def(
"fp8_scaled_mm_cpu(Tensor(a0!) mat1, Tensor(a1!) mat2, Tensor(a2!) "
@@ -468,6 +455,23 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
ops.impl("causal_conv1d_update_cpu", torch::kCPU, &causal_conv1d_update_cpu);
#endif
#if (defined(__AVX512BF16__) && defined(__AVX512F__) && \
defined(__AVX512VNNI__)) || \
defined(__riscv)
// Adapted from sglang: INT4 W4A8 kernels
ops.def(
"convert_weight_packed_scale_zp(Tensor weight, Tensor qzeros, Tensor "
"scales, int quant_method_4bit) -> (Tensor, "
"Tensor, Tensor)");
ops.impl("convert_weight_packed_scale_zp", torch::kCPU,
&convert_weight_packed_scale_zp);
ops.def(
"int4_scaled_mm_cpu(Tensor(a0!) x, Tensor(a1!) w, Tensor(a2!) w_zeros, "
"Tensor(a3!) w_scales, Tensor? bias) -> Tensor");
ops.impl("int4_scaled_mm_cpu", torch::kCPU, &int4_scaled_mm_cpu);
#endif
// Adapted from sglang: GDN kernels
ops.def(
"chunk_gated_delta_rule_cpu(Tensor query, Tensor key, Tensor value, "