Prhub

#45269 [CPU][RISC-V] Add RVV path for W4A8 INT4 GEMM

原始 PR 作者 wcynb1023 合并时间 2026-06-25 16:18 文件变更 5 提交数 10 评论 0 代码增减 +157 / -7

执行摘要

在 RISC-V CPU 上为 W4A8 GEMM 添加 RVV 向量化路径,性能提升 5 倍

为 RISC-V 平台提供 CPU W4A8 推理加速。PR 标题和 body 明确说明:此前只有 x86 AVX512 优化路径,RISC-V 用户只能使用标量路径,性能较差。本 PR 添加 RVV 后端填补该空白,并复用现有权重布局,降低集成成本。

值得精读,特别是 RISC-V 团队的成员和关注性能优化的工程师。该 PR 展示了如何以极少的代码入侵为现有调度添加新后端,并通过静态断言和回退保证健壮性。其设计模式和 RVV intrinsic 使用方式有参考价值。

讨论亮点

该 PR 未产生 review 评论,由维护者 bigPYJ1151 直接 approve。无实质讨论。

实现拆解

  1. 定义 RVV 基础类型和编译宏:在 csrc/cpu/cpu_types_riscv_defs.hpp 中新增 LMUL_64 宏(用于 8 元素向量),以及 fixed_u8x8_tfixed_i8x8_tfixed_i16x8_t 等固定长度向量类型,支持 VLEN=128/256 两种配置。在 csrc/cpu/sgl-kernels/vec.h 中检测 __riscv_v_min_vlen 并定义 CPU_CAPABILITY_RVV,然后包含上述头文件。

  2. 实现 RVV 微内核函数:在 csrc/cpu/sgl-kernels/gemm_int4.cpp 中添加 #if defined(CPU_CAPABILITY_RVV) 分支内的四个辅助函数:

    • load_uint4_as_int8_rvv:从 VNNI4 打包格式中利用 strided load 提取 8 个 int4 值,通过移位和掩码转为 int8。
    • gemm_accum_uint8_int8_rvv:将 uint8 激活量与 int8 权重执行乘加(vwmacc),更新 int32 累加器。
    • gemm_accum_uint4_rvv:组合 unpack 和 weight zero-point 减法。
    • _dequant_and_store_rvv:应用补偿、转换浮点、缩放并累加到输出。
    • _dequant_gemm_accum_rvv:主循环遍历行/组,对每个 8 列组调用上述核。
  3. 集成到现有调度:在 _dequant_gemm_accum 中添加 RVV 分支:当 defined(CPU_CAPABILITY_RVV) 且块形状符合时(N==32)调用 _dequant_gemm_accum_rvv,否则进入标量回退。

  4. Python 层开启 RVV 路径:在 vllm/model_executor/kernels/linear/mixed_precision/cpu.pyprocess_weights_after_loading 中,增加对 RISC-V 架构的检测(current_platform.get_cpu_architecture() == CpuArchEnum.RISCV),使得 layer.use_w4a8 在 AMX 或 RISC-V 时均可启用。

  5. 调整 cmake 构建选项:修正 cmake/cpu_extension.cmake 中的 RVV VLEN 检查,允许 VLLM_RVV_VLEN=0 表示强制编译标量版本,并完善注释。

文件 模块 状态 重要度
csrc/cpu/sgl-kernels/gemm_int4.cpp 内核 modified 7.33
csrc/cpu/sgl-kernels/vec.h 向量宏 modified 5.44
csrc/cpu/cpu_types_riscv_defs.hpp RVV 类型 modified 5.41
vllm/model_executor/kernels/linear/mixed_precision/cpu.py 调度层 modified 4.88
cmake/cpu_extension.cmake 构建脚本 modified 3.05

关键符号

load_uint4_as_int8_rvv gemm_accum_uint8_int8_rvv gemm_accum_uint4_rvv _dequant_and_store_rvv _dequant_gemm_accum_rvv

关键源码片段

csrc/cpu/sgl-kernels/gemm_int4.cpp core-logic

核心修改,新增 RVV 实现和调度分支

// RVV 核心实现:从 VNNI4 布局加载 8 个 int4 权重并转换为 int8
// 使用 strided load(vlse8)高效解包,然后通过移位和掩码提取低 4 位
// group 参数选择 8 列组(0-3),对应 32 列块中的子列位置
// 返回 int8x8 向量,其中每个元素是解包后的权重值template <int64_t N, int64_t ldb, int group>
inline fixed_i8x8_t load_uint4_as_int8_rvv(const uint8_t* __restrict__ B, int64_t k) {
  constexpr int64_t n_group_size = 8; // 每组 8 列
  constexpr int64_t vnni_size = 4; // VNNI 维度
  static_assert(N == 32); // 当前仅支持 32 列块
  static_assert(ldb == N / 2); // 每行步长等于一半列数
  static_assert(group >= 0 && group < N / n_group_size); // 4 个组  const int64_t ki = k % vnni_size;
  const int64_t k_base = k - ki;
  constexpr int64_t packed_group = group / 2;
  const uint8_t* packed_ptr = B + k_base * ldb + packed_group * n_group_size * vnni_size + ki;  // 跨步加载:每隔 vnni_size 个字节取一个 vnni_size 字节元素,共 n_group_size 个
  fixed_u8x8_t packed = RVVI(__riscv_vlse8_v_u8, LMUL_64)(packed_ptr, vnni_size, n_group_size);
  // 高 4 位组(group 奇数)则需要右移 4 位
  if constexpr (group % 2 == 1) {
    packed = RVVI(__riscv_vsrl_vx_u8, LMUL_64)(packed, 4, n_group_size);
  }
  fixed_u8x8_t nibbles = RVVI(__riscv_vand_vx_u8, LMUL_64)(packed, 0x0f, n_group_size);
  return RVVI4(__riscv_vreinterpret_v_u8, LMUL_64, _i8, LMUL_64)(nibbles);
}// int8 与 uint8 的乘加:先将 int8 扩展到 int16,再用 vwmacc 与 uint16 乘加
inline fixed_i32x8_t gemm_accum_uint8_int8_rvv(fixed_i32x8_t acc, uint8_t a, fixed_i8x8_t b) {
  constexpr int64_t vl = 8;
  fixed_i16x8_t b_i16 = RVVI(__riscv_vsext_vf2_i16, LMUL_128)(b, vl);
  return RVVI(__riscv_vwmacc_vx_i32, LMUL_256)(acc, static_cast<int16_t>(a), b_i16, vl);
}// 解包 + 减 zero-point + 乘加(完整的 int4 权重计算)
template <int64_t N, int64_t ldb, int group>
inline fixed_i32x8_t gemm_accum_uint4_rvv(
    fixed_i32x8_t acc,
    const uint8_t* __restrict__ B,
    const int8_t* __restrict__ qzeros_b,
    uint8_t a,
    int64_t k) {
  constexpr int64_t n_group_size = 8;
  fixed_i8x8_t b = load_uint4_as_int8_rvv<N, ldb, group>(B, k);
  // 加载对应组的 zero-point(每组 8 个 int8)
  fixed_i8x8_t qzeros =
      RVVI(__riscv_vle8_v_i8, LMUL_64)(qzeros_b + group * n_group_size, n_group_size);
  // 权重值减去 zero-point
  b = RVVI(__riscv_vsub_vv_i8, LMUL_64)(b, qzeros, n_group_size);
  return gemm_accum_uint8_int8_rvv(acc, a, b);
}
// 对 32 列块中每组 8 列进行反量化(补偿、缩放)并累加到浮点输出
// 接受 int32 累加器,消除激活 zero-point 贡献,转换为 float,乘以 scales,再与旧值相加
template <int group>
inline void _dequant_and_store_rvv(
    float* __restrict__ C,
    fixed_i32x8_t acc,
    const float* __restrict__ scales_a,
    const int32_t* __restrict__ qzeros_a,
    const float* __restrict__ scales_b,
    const int32_t* __restrict__ compensation,
    int64_t m,
    int64_t ldc) {
  constexpr int64_t n_group_size = 8;
  constexpr int64_t n = group * n_group_size;
  constexpr int64_t vl = n_group_size;  // 加载 compensation(预先计算的激活 zero-point * 权重的逐列和)
  fixed_i32x8_t comp = RVVI(__riscv_vle32_v_i32, LMUL_256)(compensation + n, vl);
  // 激活 zero-point 补偿
  fixed_i32x8_t zp_comp = RVVI(__riscv_vmul_vx_i32, LMUL_256)(comp, qzeros_a[m], vl);
  acc = RVVI(__riscv_vsub_vv_i32, LMUL_256)(acc, zp_comp, vl);  // 转 float 并乘以激活 scale 和权重 scale
  fixed_fp32x8_t acc_f = RVVI(__riscv_vfcvt_f_x_v_f32, LMUL_256)(acc, vl);
  acc_f = RVVI(__riscv_vfmul_vf_f32, LMUL_256)(acc_f, scales_a[m], vl);
  fixed_fp32x8_t scale_b = RVVI(__riscv_vle32_v_f32, LMUL_256)(scales_b + n, vl);
  acc_f = RVVI(__riscv_vfmul_vv_f32, LMUL_256)(acc_f, scale_b, vl);  // 累加到已有的浮点缓冲区(可能已包含 bias)
  float* c_ptr = C + m * ldc + n;
  fixed_fp32x8_t c_old = RVVI(__riscv_vle32_v_f32, LMUL_256)(c_ptr, vl);
  fixed_fp32x8_t c_new = RVVI(__riscv_vfadd_vv_f32, LMUL_256)(c_old, acc_f, vl);
  RVVI(__riscv_vse32_v_f32, LMUL_256)(c_ptr, c_new, vl);
}
csrc/cpu/cpu_types_riscv_defs.hpp core-logic

新增 LMUL_64 及 int8/uint8/int16 固定向量类型,为 RVV 内核提供基础类型

// VLEN 到 LMUL 的映射:LMUL_64 用于 8 元素 int8/uint8 向量
// 在 VLEN=128 时对应 mf2,VLEN=256 时对应 mf4
#if __riscv_v_min_vlen == 128
  #define LMUL_64 mf2
  #define LMUL_128 m1
  #define LMUL_256 m2
  #define LMUL_512 m4
  #define LMUL_1024 m8
  #define BOOL_256 b16
  #define BOOL_512 b8
#elif __riscv_v_min_vlen == 256
  #define LMUL_64 mf4
  #define LMUL_128 mf2
  #define LMUL_256 m1
  #define LMUL_512 m2
  #define LMUL_1024 m4
  #define BOOL_256 b32
  #define BOOL_512 b16
#else
  #error "cpu_types_riscv_defs.hpp: unsupported __riscv_v_min_vlen"
#endif// 关键固定长度向量类型定义
// uint8 / int8 (8 元素)
typedef RVVTYPE(vuint8, LMUL_64, _t) fixed_u8x8_t
    __attribute__((riscv_rvv_vector_bits(64)));
typedef RVVTYPE(vint8, LMUL_64, _t) fixed_i8x8_t
    __attribute__((riscv_rvv_vector_bits(64)));// int16 (8 元素,LMUL_128)
typedef RVVTYPE(vint16, LMUL_128, _t) fixed_i16x8_t
    __attribute__((riscv_rvv_vector_bits(128)));

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  • 兼容性风险:RVV 内核通过条件编译隔离,不影响 x86 或 ARM 等平台。标量回退存在,即使 RVV 代码编译失败也可安全降级。
  • 正确性风险:static_assert 强制 N==32,若未来引入不同块形状会导致编译错误,现有调用均满足该要求。无单元测试,依赖集成测试覆盖。
  • 性能风险:RVV 内核仅支持 VLEN=128/256,更小的 VLEN 会被标量回退兜底,不会产生错误但性能可能下降。
  • 安全风险:无。
  • 用户影响:RISC-V 用户将获得显著的推理性能提升(约 5 倍),对于 RISC-V 服务器的推理部署有重要价值。
  • 系统影响:编译系统新增 VLLM_RVV_VLEN 选项,可以为 RISC-V 选择目标 VLEN;默认识别 /proc/cpuinfo 自动选择。
  • 团队影响:需要维护 RVV 内核的持续正确性,在 CI 中缺乏 RISC-V 硬件覆盖。
  • 影响范围:受限,只修改了 CPU W4A8 算子一个路径。
缺少测试覆盖 平台特定代码 RVV kernel 仅支持 N=32

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论