CANN 鸿蒙端侧 GatherDequantInt8 自定义算子:基于 Ascend C 的 INT8 Embedding 查表与反量化融合实现
2026/9/18 16:28:46 网站建设 项目流程

CANN 鸿蒙端侧 GatherDequantInt8 自定义算子:基于 Ascend C 的 INT8 Embedding 查表与反量化融合实现

【免费下载链接】cann-recipes-harmony-infer本项目为鸿蒙开发者提供基于CANN平台的业务实践案例,方便开发者参考实现端云能力迁移及端侧推理部署。项目地址: https://gitcode.com/cann/cann-recipes-harmony-infer

导读

GatherDequantInt8是 CANN 开源仓库 cann-recipes-harmony-infer 中面向鸿蒙端侧(Kirin 9020 处理器)提供的 Ascend C 自定义算子样例:它把「INT8 embedding 查表」与「逐行(per-token)非对称反量化」融合为单个算子,直接在 AI Core 上以 uint8 读表、Cast 到 fp16 并逐行反量化,解决端侧模型 embedding 无法以 INT8 内置进图的痛点。读完本文,你将掌握该算子的数学语义、算子规格、Host/Device 两侧源码实现原理、ONNX 前端适配方式,以及从环境准备、编译安装到单算子测试的完整实战流程。

应用背景:为什么需要自定义 Gather 算子

在端侧部署带 embedding 的模型(如标点恢复、ASR 等)时,为压缩模型物理尺寸,通常希望把 embedding 权重以 INT8 形式内置进计算图。然而 Kirin 9020 工具链的框架 Gather 算子(GatherV2D)的data输入仅支持 fp32/fp16/int32,不接受 uint8/int8,导致"图内 INT8 embedding 查表"无法用标准算子组合表达。

GatherDequantInt8正是针对该场景设计的融合算子:

  • 在 AI Core 上直接以 uint8 读表、Cast 到 half、逐行反量化;
  • 使图内 embedding 可压缩到 INT8(约为 fp32 的 1/4),与外置反量化 bin 的体积对齐;
  • 数值与 fp16 反量化路径一致,在标点模型端到端部署中已验证与 fp32 embedding 的标点预测 argmax 100% 一致(见 应用场景说明)。

该样例还提供 ONNX 框架适配插件,可通过 ATC/OMG 将含该自定义节点的 ONNX 模型转换为端侧离线模型,从而在开发者自研 INT8 量化 embedding 的部署场景下,实现"图内 INT8 embedding 查表 + 反量化"。

数学语义与算子规格

数学表达式

算子对应的数学表达式为:

y[i, :] = (half(table[indices[i], :]) - zero_point[indices[i]]) * scale[indices[i]]
  • table:逐行非对称量化的 uint8 embedding 表;
  • indices:待查表的 token id;
  • scale/zero_point:逐行(per-token)的缩放因子与零点;
  • y:反量化后的 fp16 embedding。

算子规格表

名称角色数据类型形状说明
table输入uint8[V, E]逐行非对称量化的 embedding 表
indices输入int32[...]token id(按元素查表,支持任意 shape)
scale输入fp16[V]逐行 scale
zero_point输入fp16[V]逐行 zero_point
y输出fp16[..., E]反量化后的 embedding

其中 V 为词表大小,E 为 embedding 维度;输出 shape =indices.shape ++ [E],dtype 固定为 fp16。

逐行量化公式(与测试数据生成一致)

本样例test/gen_data.py采用的 per-row 非对称 uint8 量化关系为:

scale_v = (max(W_v) - min(W_v)) / 255 zp_v = -min(W_v) / scale_v q_v = round(W_v / scale_v + zp_v)

即在生成测试数据时,先对原始 fp32 embedding 逐行求 min/max 得到 scale 与 zero_point,再量化到 uint8;golden 输出则严格按 fp16 计算路径(half(q) - zp) * scale生成,用于与 device kernel 对齐验证(见 gen_data.py)。

支持的产品型号

  • Kirin 9020 处理器系列产品

如需适配 Kirin X90 / 9030,需同步修改CMakePresets.jsonASCEND_COMPUTE_UNIT与 op_host/gather_dequant_int8.cpp 中AICore().AddConfig(...)

算子工程目录结构

gather_dequant_int8_custom ├── build_and_install.sh # 编译安装脚本 ├── CMakeLists.txt ├── CMakePresets.json # 编译配置(ASCEND_CANN_PACKAGE_PATH / 算力型号) ├── framework │ ├── CMakeLists.txt │ └── onnx_plugin │ ├── CMakeLists.txt │ └── gather_dequant_int8_plugin.cpp # ONNX 前端适配插件 ├── op_host │ ├── CMakeLists.txt │ ├── gather_dequant_int8.cpp # 原型注册 / InferShape / InferDataType / Tiling │ └── gather_dequant_int8_tiling.h # TilingData 定义 ├── op_kernel │ ├── CMakeLists.txt │ └── gather_dequant_int8.cpp # 核函数实现 └── test ├── create_onnx.py # 生成单算子 ONNX 测试模型 └── gen_data.py # 生成输入与 golden 数据

源码级实现原理

该算子的实现遵循 Ascend C 自定义算子工程的标准三段式结构:Host 侧(原型注册、Shape/DType 推导、Tiling)、Device 侧(核函数)、框架适配(ONNX 插件)。

Host 侧:原型注册与 Tiling

op_host/gather_dequant_int8.cpp 承担三部分职责:

  1. 算子原型注册GatherDequantInt8类中声明 4 个 REQUIRED 输入(tableuint8 ND、indicesint32 ND、scalefp16 ND、zero_pointfp16 ND)与 1 个输出y(fp16 ND),并通过AICore().SetTiling(...)绑定 Tiling 函数、AICore().AddConfig("kirin9020")声明算力配置。

  2. InferShape:输出y的 shape 由indices的 rank 与table最后一维拼接得到,即y.shape = indices.shape ++ [E],其中E = table.shape[-1]

  3. InferDataType:输出固定为 fp16(与 scale / zero_point 一致)。

Tiling 逻辑中,从输入 shape 解析出三个关键参数并写入 TilingData:

  • vocab(V):table 第 0 维;
  • embDim(E):table 最后一维;
  • numIndices:indices 全元素乘积,即总查表次数。

同时设置SetBlockDim(1),即 Kirin 9020 AI Core 单核执行,且 workspace 大小为 0。

TilingData 结构定义于 op_host/gather_dequant_int8_tiling.h:通过BEGIN_TILING_DATA_DEF声明numIndicesembDimvocab三个 uint32 字段,并以REGISTER_TILING_DATA_CLASS完成注册,供核函数侧GET_TILING_DATA读取。

Device 侧:核函数实现

op_kernel/gather_dequant_int8.cpp 实现核函数KernelGatherDequantInt8,其设计要点:

  • 查表数据驻留 UBindices/scale/zero_point一次性载入统一缓冲区(UB),供按 idx 随机访问,通过DataCopyPad处理非 32B 对齐场景,并用PipeBarrier<PIPE_ALL>()保证 MTE2 搬运完成后标量单元GetValue才能读取。
  • 逐 token 流水:对每个 index,依次执行CopyIn(uint8 行)→Cast(uint8 转 half)→Adds(减 zero_point)→Muls(乘 scale)→CopyOut(half 行写回),其中Adds/Muls对应数学表达式中的(q - zp) * scale
  • Double buffer 并行:通过BUFFER_NUM = 2的队列配置,让搬运与计算流水并行。
  • 架构细节:由于 dav_l310(kirin9020)不允许标量 half 算术,代码将 zero_point 取负运算放到 float 上进行再转回 half;同时idxBuf/scaleBuf/zpBuf/calcBuf使用VECCALC位置,输入/输出行队列使用VECIN/VECOUT

核函数入口gather_dequant_int8通过GET_TILING_DATA读取 Host 侧下发的 tiling 数据后初始化并执行Process()

ONNX 前端适配插件

framework/onnx_plugin/gather_dequant_int8_plugin.cpp 将 ONNX 图中op_type=GatherDequantInt8的自定义节点映射到 GE 自定义算子,使omg --framework=5能解析并入图:

  • 通过REGISTER_CUSTOM_OP("GatherDequantInt8")FrameworkType(ONNX)声明框架类型;
  • OriginOpType同时覆盖裸算子名与多种 domain 前缀(如custom::GatherDequantInt8ai.onnx::1::GatherDequantInt8直至ai.onnx::16::GatherDequantInt8);
  • 该节点无属性,输入/输出顺序与 GE 原型一致,故直接使用AutoMappingByOpFn自动映射。

操作步骤

1. 环境准备

参考 环境准备 完成环境搭建,核心前提如下:

  • Python >= 3.7.0、gcc >= 7.3.0、cmake >= 3.16.0,建议使用 Ubuntu 22.04 以上环境(依赖 glibc 2.34+);
  • 安装鸿蒙社区版 CANN 开发套件包Ascend-cann-toolkit_${cann_version}_linux-${arch}-mobile-station.run,安装命令形如:
chmod +x Ascend-cann-toolkit_${cann_version}_linux-${arch}-mobile-station.run ./Ascend-cann-toolkit_${cann_version}_linux-${arch}-mobile-station.run --install --force --install-path=${install_path}
  • 配置环境变量:
source /usr/local/Ascend/cann-${cann_version}/set_env.sh

编译前确认 CMakePresets.json 中ASCEND_CANN_PACKAGE_PATH指向正确的 toolkit 安装路径(一般为${install_path}/cann)。该文件还集中定义了ASCEND_COMPUTE_UNIT(当前为kirin9020)、vendor_namecustomize)、ENABLE_TESTENABLE_CROSS_COMPILE等编译选项,适配新算力型号时主要修改ASCEND_COMPUTE_UNIT一项。

2. 编译安装

在算子工程目录下执行:

chmod +x build_and_install.sh ./build_and_install.sh

build_and_install.sh 内部会依次完成:设置ASCEND_HOME_PATHsetenv.bash环境、将CMakePresets.json中的默认 CANN 路径替换为实际安装路径、以defaultpreset 配置并编译binarypackage目标,最后执行生成的custom_opp_${OS_ID}_${arch}.run --quiet完成安装。编译产物为自定义算子 run 包并自动安装到packages/vendors/customize/下。详细流程还可参考 算子工程编译安装指南。

3. 单算子测试

cd test python3 create_onnx.py # 生成 GatherDequantInt8.onnx python3 gen_data.py # 生成 table/indices/scale/zero_point.bin 与 golden output.bin
  • create_onnx.py 使用默认规格 V=1024、E=256、N=30(SEQ_LEN)构造含单自定义节点GatherDequantInt8的 ONNX 模型,输入/输出类型与顺序和 GE 原型保持一致(table uint8、indices int32、scale/zero_point fp16、y fp16),opset 版本设为 11。
  • gen_data.py 以固定随机种子生成量化后的 table、随机 indices、逐行 scale/zero_point,并按 fp16 计算路径生成 goldenoutput.bin,供精度比对使用。

随后可通过 ATC 工具转换测试模型(参考 ATC 工具使用指南),调用鸿蒙维测接口完成单算子的精度与性能验证。

总结

GatherDequantInt8以"一个自定义算子替代标准算子无法表达的 INT8 embedding 查表 + 反量化"为切入点,完整展示了鸿蒙端侧 Ascend C 算子开发的通用范式:Host 侧原型注册与 Tiling 下发、Device 侧 double-buffer 流水核函数、ONNX 前端插件适配,以及配套的单算子 ONNX 测试与 golden 数据生成链路。对需要在 Kirin 9020 上以 INT8 压缩 embedding 并保持 fp16 精度路径一致的开发者而言,本样例既是一份可直接复用的算子实现,也是学习端侧自定义算子全流程的最佳参考。

【免费下载链接】cann-recipes-harmony-infer本项目为鸿蒙开发者提供基于CANN平台的业务实践案例,方便开发者参考实现端云能力迁移及端侧推理部署。项目地址: https://gitcode.com/cann/cann-recipes-harmony-infer

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询