Skip to content

[DEV] 为 mul_scalar / moe_sum / prepare_moe_input / rwkv5_wkv / mamba_selective_scan / dequantize_gptq 补齐 Metax 后端 #1395

Description

@Lfan-ke

目标版本

main

功能描述

这六个算子目前只有 nvidia/ 实现,没有 metax/,因此在 Metax GPU 上不可用:

算子 说明
mul_scalar 逐元素乘标量(走 elementwise 框架)
moe_sum MoE top-k 专家输出求和
prepare_moe_input MoE 专家偏移 / 置换 / 概率整理
rwkv5_wkv RWKV-5 的 WKV 递推
mamba_selective_scan Mamba 选择性扫描
dequantize_gptq GPTQ 4-bit 权重反量化

它们都没有 cuDNN / cuBLAS / CUTLASS / 内联 PTX 依赖,可以按同族算子已经确立的机械移植路径补齐:运行时头(mc_runtime.h / hc_runtime.h)、devices/metax/*、命名空间、cudaStream_thcStream_tdevice::nvidia::Handledevice::metax::Handle,设备端 kernel 原样复用;再在 operator.cc 的四处派发(create / get-workspace / calculate / destroy)接入 INFINI_DEVICE_METAX

实现中需要注意的几点(都是在 MetaX C500 上真机编译才暴露出来的)

  1. moe_sum / prepare_moe_input / dequantize_gptq 的 nvidia 实现整体包在 #ifdef ENABLE_NVIDIA_API 里。照搬到 .maca 会被预处理器整个掏空 —— 目标文件照样生成,但链接时才发现符号全没了(undefined symbol: op::moe_sum::metax::Descriptor::~Descriptor())。
  2. devices/metax/metax_common.h 会引入 metax_ht2mc.h,那里才有 hcStream_tmcStream_thpcc_bfloat16maca_bfloat16 这一层兼容宏。所以它必须排在 metax_kernel_common.h 之前,否则在 --use-mc=y 下这些类型全都不认识。
  3. __nv_bfloat16 的别名在 metax_kernel_common.h 里,用到 bf16 的算子必须包含它。
  4. CHECK_CUDA 在 metax 侧叫 CHECK_METAXmetax_kernel_common.h:23)。

验证

在 MetaX C500(MACA 3.5.3.20 / 驱动 3.8.30)上 XMAKE_ROOT=y xmake f --metax-gpu=y --use-mc=y -y && xmake && xmake install 全绿(.maca 是带 -Werror 编的),并跑 python test/infiniop/dequantize_gptq.py --metax 通过。

其余五个算子在 test/infiniop/ 下还没有测试文件,我会在 PR 里说明这一点,并且乐意补上测试。

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions