Prhub

#27947 [AMD] Fix jit-kernel-unit-test-amd: activation.cuh ROCm build + per_token CUDA-only (R165)

原始 PR 作者 michaelzhang-ai 合并时间 2026-06-16 14:44 文件变更 3 提交数 10 评论 18 代码增减 +18 / -4

执行摘要

修复 AMD JIT 内核编译与测试注册

AMD 的 jit-kernel-unit-test-amd 套件(R165)因两个测试失败而持续红色:test_activation(编译错误)和 test_per_token_group_quant_8bit(运行时阻塞)。根本原因在于 activation.cuh 的 unary kernel 使用 decltype(auto) 作为返回类型,clang-HIP 无法初始化 const 限定的函数指针;而 per_token 测试依赖的 fp8 dtype trait 在 ROCm 下未实现,且其 warp reduce 使用 32 位掩码与 HIP 64 位 warp 不兼容。

建议所有参与跨平台内核开发的团队精读该 PR。它不仅解决了一个实际的编译问题,还展示了如何在不污染 CUDA 代码库的前提下适配 AMD ROCm 平台。关键设计决策(如使用非 #ifdef 的类型别名统一返回类型、分离而非替换 CUDA reduce 实现)值得作为最佳实践推广。

讨论亮点

1. activation.cuh select_unary_kernel 实现方式

  • BBufDarkSharpness 反对在 #ifdef USE_ROCM 中复制完整函数体,建议只保留一个非 trailing 返回类型的通用版本。
  • 作者接受并在 commit 3a41983 中移除 #ifdef 分支,使用显式 unary_kernel_fn_t 返回类型,使 CUDA 和 ROCm 共享同一实现,且通过 clang-format 检查。

2. per_token GROUP reduce 处理

  • HaiShaw 建议 refine 实现,移除未使用的 tid 参数。
  • BBuf 建议保留 CUDA 实现不变,仅添加 AMD 独立分支,而不是统一修改。作者最初尝试统一为 portable warp reduce(commit 739cecf),但最终在 f7ffafa 中采纳 BBuf 意见,CUDA 路径恢复到与 main 字节一致,AMD 路径在 #ifdef USE_ROCM 下分离。

3. 移除 per_token 测试的 AMD CI 注册

  • HaiShaw 询问“remove amd test?”,作者解释该测试因框架级 fp8 阻塞从未在 AMD 上通过,且注册是意外引入(源自未合并的 #27717),移除不会降低实际覆盖。

实现拆解

  1. 修复 activation.cuh 在 ROCm 下的编译:在 python/sglang/jit_kernel/csrc/elementwise/activation.cuh 中,为 ActivationKernel 添加 unary_kernel_fn_t = decltype(&act_kernel<T, kReLU2, kUsePDL>) 类型别名,并将 select_unary_kernel 的返回类型从 trailing decltype(unary_kernel<...>) 改为直接使用 unary_kernel_fn_t。这样避免了 const 自动推导导致的初始化问题,且与二进制路径的现有模式一致,CUDA 编译完全不变。
  2. 添加 AMD 专用的 GroupReduceMax 实现:在 python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh 中,用 #ifdef USE_ROCM 条件编译新增基于 device::warp::reduce_max<kThreadsPerGroup> 的版本。CUDA 原有的 __shfl_xor_sync 实现保持字节级一致,仅移除未使用的 tid 参数。HIP 的 64 位 warp 需要显式子组宽度,故不可复用原 32 位掩码逻辑。
  3. 调整 per_token 测试的 AMD CI 注册:在 test/registered/jit/test_per_token_group_quant_8bit.py 中,移除 register_amd_ci 的导入和调用,使该测试仅用于 CUDA。原因是框架层 fp8_e4m3_t__HIPCC__ 下缺少 dtype trait 定义,导致 TensorMatcher 无法识别 e4m3fn 输出类型,无法在 AMD 上运行。CUDA 全量 1872 个配置的矩阵测试保持不变并在 nightly 中继续执行。
文件 模块 状态 重要度
python/sglang/jit_kernel/csrc/elementwise/activation.cuh JIT 内核 modified 4.0
python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh JIT 内核 modified 4.03
test/registered/jit/test_per_token_group_quant_8bit.py 测试 modified 3.52

关键符号

select_unary_kernel GroupReduceMax

关键源码片段

python/sglang/jit_kernel/csrc/elementwise/activation.cuh core-logic

核心修复:解决 ROCm 下 unary kernel 编译错误,通过添加 `unary_kernel_fn_t` 类型别名并修改 `select_unary_kernel` 返回类型。

struct ActivationKernel {
    // ...
    using kernel_fn_t = decltype(&act_and_mul_kernel<T, ActivationKind::kSiLU, kUsePDL, false>);
    // ADDED: concrete typedef for unary kernel function pointer (mirrors binary path)
    using unary_kernel_fn_t = decltype(&act_kernel<T, ActivationKind::kReLU2, kUsePDL>);    template <ActivationKind kAct>
    static constexpr auto unary_kernel = act_kernel<T, kAct, kUsePDL>;    // BEFORE: static auto select_unary_kernel(const std::string& type)
    // -> decltype(ActivationKernel::template unary_kernel<ActivationKind::kReLU2>)
    // AFTER: explicit non-const fn-ptr type; works for both nvcc and clang-HIP
    static unary_kernel_fn_t select_unary_kernel(const std::string& type) {
        using namespace host;
        if (type == "relu2") {
            return ActivationKernel::template unary_kernel<ActivationKind::kReLU2>;
        }
        // ... other cases
        return nullptr;
    }
};
python/sglang/jit_kernel/csrc/gemm/per_token_group_quant_8bit.cuh core-logic

添加 AMD 专用的 GroupReduceMax 实现,保留 CUDA 路径不变。

namespace {
constexpr int kThreadsPerGroup = 16;#ifdef USE_ROCM
// AMD implementation: HIP warps are 64-wide and require explicit sub-group width.
// Delegate to the portable warp-reduce primitive, which emits __shfl_xor with
// the correct kThreadsPerGroup sub-group width.
__device__ __forceinline__ float GroupReduceMax(float val, const int /*tid*/) {
    return device::warp::reduce_max<kThreadsPerGroup>(val);
}
#else
// CUDA implementation (unchanged from main): hand-rolled reduction using
// 32-bit shuffle masks; assumes warp width <= 32.
__device__ __forceinline__ float GroupReduceMax(float val, const int tid) {
    unsigned mask = threadIdx.x % 32 >= 16 ? 0xffff0000 : 0x0000ffff;
    val = fmaxf(val, __shfl_xor_sync(mask, val, 8));
    val = fmaxf(val, __shfl_xor_sync(mask, val, 4));
    val = fmaxf(val, __shfl_xor_sync(mask, val, 2));
    val = fmaxf(val, __shfl_xor_sync(mask, val, 1));
    return val;
}
#endif
} // namespace

评论区精华

activation.cuh select_unary_kernel 实现方式 设计

BBuf 和 DarkSharpness 反对在 #ifdef USE_ROCM 中复制整个函数体,建议只保留一个非 trailing 返回类型的通用版本。

结论:作者在 commit 3a41983 中移除 #ifdef 分支,使用显式 unary_kernel_fn_t 返回类型,两后端共享同一实现。 · 已解决

per_token GroupReduceMax 处理 设计

BBuf 建议保留 CUDA 实现并新增 AMD 独立分支,而不是统一修改。HaiShaw 建议 refine。

结论:作者最终在 f7ffafa 中恢复 CUDA 路径为与 main 字节一致,AMD 分支在 #ifdef USE_ROCM 下使用 device::warp::reduce_max。 · 已解决

移除 per_token 测试的 AMD CI 注册 测试

HaiShaw 询问为何移除 AMD 测试注册。作者解释该测试因框架级 fp8 阻塞从未在 AMD 上通过,且注册是意外引入。

结论:移除被接受,确认不会降低实际覆盖,且后续 AMD fp8 支持上线后可重新注册。 · 已解决

风险与影响

  • CUDA 代码生成风险:activation.cuh 返回类型从 decltype(auto) 变为显式非 const 函数指针,nvcc 接受且编译后语义等价。per_token 的 CUDA GroupReduceMax 保持字节一致,无风险。
  • AMD 特有实现正确性:新增的 device::warp::reduce_max<kThreadsPerGroup> 仅在 AMD 上编译和使用,但未在 AMD CI 中直接测试(因 per_token 测试已移除注册)。若未来启用在 AMD 上运行该测试,可能暴露逻辑差异。
  • 测试覆盖下降:per_token 测试从 AMD CI 移除,导致 AMD 缺少对该量化内核的端到端验证。但由于它从未成功运行过,实际覆盖未损失,但可能掩盖后续引入的回归。
  • 维护负担:per_token 内核现在有两个实现路径(CUDA/ROCm),未来若修改 CUDA 的 reduce 逻辑,需要同步考虑 AMD 分支。
  • 用户:AMD 开发者将看到更干净的 CI 结果,减少误报。activation 测试现在可以正确编译和运行,增强对内核改动的信心。
  • 系统:CUDA 路径无变化,性能无影响。AMD 上 activation 内核通过 73 个测试。per_token 测试不会在 AMD CI 中执行。
  • 团队:修复了持续集成的中断,后续 AMD fp8 支持工作可以在此基础上进行。BBuf 和 DarkSharpness 的代码审阅意见被采纳,提高了代码质量。
AMD warp reduce 未在 CI 中直接测试 per_token 测试 AMD 注册移除可能掩盖回归 CUDA 代码生成不变,风险低

关联 Issue

#28323 [AMD] fix(jit): make activation unary kernel compile on ROCm (fix jit-kernel-unit-test-amd)

完整报告

参与讨论