Prhub

#2274 [ROCm] Support the INT4 QAT kernel on ROCm

原始 PR 作者 LZ-QWQ 合并时间 2026-08-16 17:44 文件变更 2 提交数 1 评论 0 代码增减 +39 / -14

执行摘要

为 ROCm 构建 INT4 QAT 内核并修复 hipcc 兼容性

PR 描述指出 int4_qat 内核在 ROCm PyTorch 安装下无法构建,导致 INT4 QAT 路径在 AMD GPU 上不可用。需要修复两个阻塞点:HIP warp-shuffle 掩码和 hipcc 下的 const_data_ptr 符号 mangling 不匹配。

值得精读。该 PR 展示了如何通过条件编译适配不同 GPU 编译器(nvcc vs hipcc),但改动量较小。关注点在于 ROCm 构建配置的通用性和未来维护成本。

讨论亮点

无 review 评论,因此没有公开的讨论记录。

实现拆解

  1. 修改 fake_int4_quant_cuda.cu:为核心内核函数 launch_int4_quant_kernel 添加 HIP 平台条件编译,将 __shfl_xor_sync 替换为无掩码的 __shfl_xor(适用于 32 线程组),并处理 const_data_ptr 模板在 hipcc 下的链接问题。
  2. 修改 setup.py:检测 ROCm 环境(torch.version.hip),设置 PYTORCH_ROCM_ARCH 默认 gfx950,并在非 ROCm 时仅传递 nvcc 专用标志(如 -gencode--expt-relaxed-constexpr),避免 hipcc 拒绝。
  3. 验证:在 MI355X / ROCm 7.2 上验证 144 个 case 与 CUDA 逐位一致,并通过 E2E 转换工具验证 Qwen3-0.6B 的 INT4 转换质量。
文件 模块 状态 重要度
slime/backends/megatron_utils/kernels/int4_qat/setup.py 后端 modified 6.47
slime/backends/megatron_utils/kernels/int4_qat/fake_int4_quant_cuda.cu 后端 modified 3.88

关键符号

warpReduceMax warpReduceMin fake_int4_quant_cuda

关键源码片段

slime/backends/megatron_utils/kernels/int4_qat/fake_int4_quant_cuda.cu core-logic

内核源码的 HIP 兼容性修复,是功能实现的直接载体。

// slime/backends/megatron_utils/kernels/int4_qat/fake_int4_quant_cuda.cu
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>// HIP 的 __shfl_xor_sync 要求 64 位掩码(wavefront 为 64 条 lane),0xFFFFFFFF 无法编译。
// 归约只在 32 条 lane 的组内进行,因此使用无掩码的 __shfl_xor 并指定宽度 32 即可。
#if defined(__HIP_PLATFORM_AMD__)
#define WARP_XOR(val, mask) __shfl_xor((val), (mask), 32)
#else
#define FINAL_MASK 0xFFFFFFFF
#define WARP_XOR(val, mask) __shfl_xor_sync(FINAL_MASK, (val), (mask), 32)
#endif// 设备端归约函数:使用 WARP_XOR 替代 __shfl_xor_sync(原代码中 mask 从 16 递减到 1)
__device__ __forceinline__ float warpReduceMax(float val) {
    #pragma unroll
    for (int mask = 16; mask > 0; mask >>= 1)
        val = fmaxf(val, WARP_XOR(val, mask));
    return val;
}__device__ __forceinline__ float warpReduceMin(float val) {
    #pragma unroll
    for (int mask = 16; mask > 0; mask >>= 1)
        val = fminf(val, WARP_XOR(val, mask));
    return val;
}// 在量化内核启动处,针对 HIP 平台使用静态转换以解决 clang 与 gcc 的模板 mangling 差异
void fake_int4_quant_cuda(...) {
    ...
    launch_int4_quant_kernel<scalar_t>(
#if defined(__HIP_PLATFORM_AMD__)
        // templated const_data_ptr<T> 在 hipcc 下无法链接:clang 对 enable_if 的 mangling 与构建 libtorch 的 gcc 不同
        static_cast<const scalar_t*>(x.const_data_ptr()),
#else
        x.const_data_ptr<scalar_t>(),
#endif
        out.data_ptr<scalar_t>(),
        out_scale.data_ptr<scalar_t>(),
        out_zero.data_ptr<scalar_t>(),
        ...);
}

评论区精华

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

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

风险与影响

  1. CUDA 路径风险低:所有改动均在 #if defined(__HIP_PLATFORM_AMD__)torch.version.hip is not None 保护内,确保 CUDA 编译路径不变。
  2. ROCm 路径风险:依赖 PYTORCH_ROCM_ARCH 环境变量,默认 gfx950 可能不适用于所有 AMD GPU,用户需显式设置。
  3. API 兼容性const_data_ptr 的静态转换方式可能影响未覆盖的调用点,需确认 fake_int4_quant_cuda 的所有调用路径。

影响范围:直接影响 INT4 QAT 内核的构建和运行,主要涉及 ROCm 用户(AMD GPU 集群)。对 CUDA 用户无影响。
团队影响:扩大了 Slime 在 AMD 硬件上的可用性,可能需要更新文档或 CI 以覆盖 ROCm 构建测试。

核心路径变更 缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论