CANN ops-nn GatherElementsV3 算子详解:在 NPU 上按维度聚集张量元素
【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn
导读
GatherElementsV3 是 CANN ops-nn 神经网络算子库(experimental/index/gather_elements_v3)中提供的数据聚集算子,它接收一个输入张量x和一个索引张量index,沿指定维度dim按索引位置取出元素构成输出张量y,是 embedding 检索、注意力掩码、排序重排等场景中的基础算子。本文以该算子目录下的 README 为骨架,结合其算子注册、形状推导、Tiling 计算与 NPU 内核实现的完整源码链路,讲解 GatherElementsV3 的数学定义、参数约束、aclnn 调用方式与底层工作原理,帮助你掌握在 Atlas 系列产品上使用与理解该算子的完整方法。
产品支持情况
根据 README 的说明,GatherElementsV3 的产品支持情况如下:
| 产品 | 是否支持 |
|---|---|
| Atlas A2 训练系列产品 / Atlas 800I A2 推理产品 | √ |
在源码层面,这一支持关系体现在 gather_elements_v3_def.cpp 中通过this->AICore().AddConfig("ascend910b", aicoreConfig)为ascend910b平台(即 Atlas A2 系列对应的 AI Core 架构)注册了算子的 AICore 配置,并在 gather_elements_v3_binary.json 中按不同数据类型声明了对应的算子二进制文件(GatherElementsV3_fp32、GatherElementsV3_fp16、GatherElementsV3_bf16、GatherElementsV3_int32)。
功能说明与数学定义
算子功能
GatherElementsV3 的功能是:对输入张量x中指定的维度dim进行数据聚集。索引张量index的每个位置给出一个在dim维度上的取值,用于从x中取出对应元素,最终输出张量y的形状与index相同。
计算公式
给定张量 $x$、维度 $d$ 和索引张量 $index$,定义 $n$ 是 $x$ 的维度数,$i_d$ 表示维度 $d$ 的索引,$index_{i_d}$ 表示索引张量 $index$ 在维度 $d$ 上的第 $i_d$ 个索引值。沿指定维度 $d$ 的 gather 功能可以用如下数学公式表示:
$$ gather(X,index,d){i_0,i_1,\cdots,i{d-1},i_{d+1},\cdots,i_{n-1}} = x_{i_0,i_1,\cdots,i_{d-1},index_{i_d},i_{d+1},\cdots,i_{n-1}} $$
通俗理解:输出张量在某个位置上的取值,等于输入张量在"除dim维外其余下标与输出位置一致、dim维下标取自index对应位置"处的元素。
计算示例
假设输入张量 $x=\begin{bmatrix}1 & 2 & 3\ 4 & 5 & 6\ 7 & 8 & 9\end{bmatrix}$,索引张量 $index=\begin{bmatrix}0 & 2\ 1 & 0\end{bmatrix}$,$dim = 0$,那么输出张量 $y=\begin{bmatrix}1 & 8\ 4 & 2\end{bmatrix}$,具体计算过程如下:
$$ \begin{aligned} y_{0,0}&=x_{index_{0,0}, 0}=x_{0,0}=1 \ y_{0,1}&=x_{index_{0,1}, 1}=x_{2,1}=8 \ y_{1,0}&=x_{index_{1,0}, 0}=x_{1,0}=4 \ y_{1,1}&=x_{index_{1,1}, 1}=x_{0,1}=2 \end{aligned} $$
可以看到,当dim=0时,输出元素由"index在该位置提供的行号 + 输出位置自身的列号"从x中定位取值。同理,若dim=1,则行号取自输出位置自身、列号取自index对应位置的值。
参数说明
GatherElementsV3 共包含 4 个参数,其中 3 个张量参数、1 个标量属性,汇总如下:
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x | 输入 | 公式中的 x,即被聚集的源张量 | FLOAT、FLOAT16、BFLOAT16、INT32 | ND |
| index | 输入 | 公式中的 index,提供各位置在 dim 维上的取值 | INT32 | ND |
| y | 输出 | 公式中的 y,即聚集结果张量 | FLOAT、FLOAT16、BFLOAT16、INT32 | ND |
| dim | 可选属性 | 公式中的 d,指定聚集发生的维度;默认值为 0 | Int | - |
上述定义与算子注册文件 gather_elements_v3_def.cpp 完全一致:输入x支持DT_BF16 / DT_FLOAT16 / DT_FLOAT / DT_INT32四种类型,输入index固定为DT_INT32,输出y与x类型保持一致,属性dim通过this->Attr("dim").Int(0)声明且默认值为 0,所有张量均使用FORMAT_ND格式。相应地,gather_elements_v3_binary.json 中为 float32、float16、bfloat16、int32 四种数据类型分别注册了独立的算子二进制,形状均以-2(动态 rank)声明。
关于输出形状
数学定义要求输出形状与index一致。这一点在形状推导实现 gather_elements_v3_infershape.cpp 中体现:InferShapeGatherElementsV3取context->GetInputShape(1)(即index的形状)并直接赋值给输出y(*yShape = *xShape),从而保证运行时框架能够正确推导出输出张量的形状。
约束说明
根据 README,本算子无额外约束说明。不过从实现细节可以推断出两条使用上的注意点:
dim取值必须小于x的维度数:Tiling 阶段在 gather_elements_v3_tiling.cpp 中显式校验dim >= xShape.GetDimNum()并返回GRAPH_FAILED,因此调用时传入的dim越界会直接导致算子执行失败;index支持负索引:内核在 gather_elements_v3.h 中对indexVal < 0的情况执行indexVal += xGatherDim_,将负索引换算为合法的非负下标(与 PyTorch 等框架的语义一致),换算后仍然越界的索引由硬件内存访问保护兜底。
调用说明:基于 aclnn 的两阶段接口
README 给出了一种调用方式——aclnn 调用,样例代码位于 test_aclnn_gather_elements_v3.cpp。该样例完整展示了 CANN 算子的标准调用范式,整体流程可分为以下步骤:
1. 环境初始化(固定写法)
int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); // 内部依次调用 aclInit、aclrtSetDevice、aclrtCreateStreamInit函数依次执行aclInit(nullptr)初始化 ACL 运行时、aclrtSetDevice(deviceId)绑定设备、aclrtCreateStream(stream)创建任务流,是所有 CANN 应用启动时的固定前置步骤。
2. 构造输入输出 Tensor
样例以selfShape = {4, 2}、indexShape = {4, 2}、outShape = {4, 2}为例:
std::vector<float> selfHostData = {0, 1, 2, 3, 4, 5, 6, 7}; std::vector<int32_t> indexHostData = {0, 1, 2, 3, 3, 2, 1, 0}; int32_t dim = 0;CreateAclTensor辅助函数完成三件事:aclrtMalloc申请 device 侧内存、aclrtMemcpy将 host 数据拷贝到 device、aclCreateTensor基于 shape/strides/ND 格式创建aclTensor句柄。其中x使用ACL_FLOAT、index使用ACL_INT32,与参数表中声明的数据类型一一对应。
3. 两阶段算子调用
aclnn 接口采用"先算 workspace 再执行"的两阶段设计:
uint64_t workspaceSize = 0; aclOpExecutor* executor; // 第一段:计算 workspace 大小并获取执行器 ret = aclnnGatherElementsV3GetWorkspaceSize(self, index, dim, out, &workspaceSize, &executor); // 根据计算出的 workspaceSize 申请 device 内存 if (workspaceSize > 0) { ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 第二段:真正下发算子任务 ret = aclnnGatherElementsV3(workspaceAddr, workspaceSize, executor, stream);第一段接口aclnnGatherElementsV3GetWorkspaceSize入参依次为self(即 x)、index、dim、out,输出 workspace 大小与执行器句柄;第二段接口aclnnGatherElementsV3负责将任务提交到指定 stream 执行。workspace 是算子运行时需要的临时内存,由调用方按接口计算出的尺寸自行申请与释放。
4. 同步、取回结果与资源释放
ret = aclrtSynchronizeStream(stream); // 同步等待任务执行结束 ret = aclrtMemcpy(resultData.data(), ..., outDeviceAddr, ..., ACL_MEMCPY_DEVICE_TO_HOST); // 取回结果 // 释放:aclDestroyTensor ×3、aclrtFree ×4(含 workspace)、aclrtDestroyStream、aclrtResetDevice、aclFinalize样例最终会逐元素打印result[i]。以示例数据(x为 0~7 的 4×2 张量、dim=0)可以验证:y[i][j] = x[index[i][j]][j],即输出应为{0, 5, 4, 7, 6, 5, 2, 1}。
源码级实现原理
算子定义注册(OpDef)
gather_elements_v3_def.cpp 中通过OP_ADD(GatherElementsV3)完成算子注册,除了声明输入/输出/属性外,还通过OpAICoreConfig指定了平台相关能力:
DynamicCompileStaticFlag(true):支持编译期静态信息参与动态编译;DynamicRankSupportFlag(true):支持动态 rank,即x/index的维度数可在运行期变化;DynamicShapeSupportFlag(true):支持动态 shape;PrecisionReduceFlag(true):允许在满足精度要求的前提下进行精度降低优化。
这些标志与二进制配置文件中"shape": [-2]的动态形状声明相呼应,共同支撑算子的动态 shape 能力。
Tiling 计算(Host 侧)
Tiling 是算子性能的关键,其职责是把大张量切分成适合 AI Core 上 UB(Unified Buffer)的 Tile。核心逻辑在 gather_elements_v3_tiling.cpp,主要步骤包括:
- 平台信息获取:通过
PlatformAscendC获取 UB 大小ubSize与可用核数coreNum; - 维度展平(flatten):以
dim为界,把x与index的 shape 分别拆成Pre / Gather / Post三段乘积,例如xPre = xShape[0..dim-1]、xGather = xShape[dim]、xPost = xShape[dim+1..]。展平后,一次聚集行的跨度即为xPostDim; - 核数规划(usedCores):按"总行数 = idxPre × idxGather"计算,若总行数小于核数则只启用对应数量的核,避免空转;
- Tile 大小计算:预留
RESERVED_UB = 1024字节、按双缓冲折半(/2)得到可用 UB,再结合"每元素需要sizeof(int32_t) + typeLength字节(index 与 x 各一份)"计算单次可处理的最大元素数,并按 32 字节对齐约束向下取整,最后不超过idxPost; - 结果写入:将上述参数写入
GatherElementsV3TilingData结构体,并通过context->SetBlockDim(usedCores)设定内核启动的核数。
Tiling 数据结构定义在 gather_elements_v3_tiling_data.h,共 8 个uint32_t字段(xPreDim/xGatherDim/xPostDim/idxPreDim/idxGatherDim/idxPostDim/usedCores/tileSize),由 Host 侧填充、Kernel 侧读取。
内核实现(Device 侧)
内核入口 gather_elements_v3.cpp 是一个带模板参数schMode(调度模式)的__global__ __aicore__函数,模板参数通过 gather_elements_v3_tiling_key.h 中ASCENDC_TPL_ARGS_DECL声明的 0/1 两种调度模式实例化。入口函数注册并读取 Tiling 数据后,实例化NsGatherElementsV3::GatherElementsV3<DTYPE_X>并依次调用Init与Process。
算子类的核心实现在 gather_elements_v3.h,要点包括:
- 流水结构:使用
TPipe配合TQue<TPosition::VECIN, BUFFER_NUM>(index 输入队列)与TQue<TPosition::VECOUT, BUFFER_NUM>(y 输出队列),BUFFER_NUM = 2实现双缓冲,CopyIn → Compute → CopyOut三级流水重叠; - 数据搬入:
CopyIn通过DataCopyPad按字节数拷贝 index 片段到局部内存,并对齐到ALIGN_BYTES = 32字节; - 核心计算:
Compute逐元素处理,先做负索引修正(indexVal < 0时indexVal += xGatherDim_),再按展平后的偏移公式xRealOffset = xBase + indexVal * xPostDim_ + (postStart + i)从全局内存xGm_取数写入yLocal——这正是 README 数学公式的代码化表达; - 任务切分:
Process按rowId = coreId_; rowId < totalRows; rowId += coreNum_的方式在多个核间按行轮转,每行内部再按tileSize分片遍历idxPostDim,兼顾了多核并行与 UB 容量限制。
构建集成
算子目录的 CMakeLists.txt 采用通用的子目录聚合模式:默认遍历并添加各子目录(op_host、op_kernel等),当未开启ENABLE_TEST且未开启BENCHMARK时排除tests目录,避免测试代码进入发布构建。整体算子库通过仓库根目录 CMakeLists.txt 统一组织编译。
总结
GatherElementsV3 是 CANN ops-nn 中一个"接口简单、实现讲究"的数据聚集算子:对外,它通过aclnnGatherElementsV3GetWorkspaceSize+aclnnGatherElementsV3两阶段接口向用户提供与 PyTorchgather语义一致的按维度取数能力;对内,它借助算子注册、动态 shape 支持、基于 UB 容量的 Tiling 切分、多核按行轮转以及双缓冲流水,将索引寻址的访存密集型计算高效地映射到 Atlas A2 系列产品的 AI Core 上。理解其参数约束、负索引语义与两阶段调用范式,即可在自研模型中安全、高效地使用该算子,也可为阅读同类 aclnn 算子的源码实现提供参考范本。
【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考