CANN ops-nn 分段求和算子 SegmentSum 详解:原理、约束与 GE 图模式调用实战
2026/9/23 22:26:29 网站建设 项目流程
  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-nn
点击查看免费下载

SegmentSum(分段求和)是 CANN ops-nn 算子库中的索引类计算算子,按分段索引对输入 Tensor 的若干行进行分组求和。本文以 index/segment_sum/README.md 为核心骨架,结合仓库内算子定义、Tiling、Kernel、配置与测试源码,系统讲解 SegmentSum 的功能语义、参数与约束、底层实现机制,并通过完整示例演示如何在 GE 图模式下调用该算子。

产品支持情况

SegmentSum 在当前仓库中的产品适配情况如下表所示:

产品是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品×
Atlas A2 训练系列产品/Atlas A2 推理系列产品×
Atlas 200I/500 A2 推理产品×
Atlas 推理系列产品×
Atlas 训练系列产品×

当前仅有 Ascend 950 系列(950PR/950DT)支持 SegmentSum,其余 Atlas 训练/推理产品暂不支持。这一结论同样可以从算子注册代码中得到印证:op_host/segment_sum_def.cpp 中仅对ascend950ascend350两个 SoC 配置注册了 AICore 配置,且 tiling 与 kernel 实现均位于arch35目录下(arch35对应 Ascend 950 的架构代号)。同时 op_host/config/ 下只提供了ascend950ascend350两套算子二进制 JSON 配置,与 README 的支持矩阵保持一致。

功能说明

SegmentSum 的功能是:对输入 Tensorx按分段索引segment_ids进行求和。其数学定义如下:

$$ y[i] = \sum_{\substack{j\ \text{segment_ids}[j] = i}} x[j] $$

即:遍历所有满足segment_ids[j] == i的索引j,将对应的x[j]逐元素累加到y[i];若某个段i没有任何元素属于它,则y[i] = 0。该算子与 TensorFlow 的tf.math.segment_sum语义兼容,这一点在 op_graph/segment_sum_proto.h 的注释中有明确说明,测试脚本 tests/assets/golden.py 也直接以tf.math.segment_sum的结果作为 Golden 基准进行对拍验证。

计算示例

README 中给出的标准用例为:

  • 输入 Tensor

$$ x = \begin{bmatrix} [1 & 2] \ [3 & 4] \ [5 & 6] \ [7 & 8] \end{bmatrix} $$

  • 分段索引 Tensor

$$ segment_ids = [0, 0, 1, 2] $$

  • 输出 Tensor

$$ y = \begin{bmatrix} [4 & 6] \ [5 & 6] \ [7 & 8] \end{bmatrix} $$

推导过程:segment_ids中索引 0、1 对应段 0,因此y[0] = x[0] + x[1] = [1+3, 2+4] = [4, 6];索引 2 对应段 1,y[1] = x[2] = [5, 6];索引 3 对应段 2,y[2] = x[3] = [7, 8]。由于每个元素按行(第 0 维)分组,SegmentSum 的求和粒度是“整行累加”,即逐列分别相加。

该示例在 GE 图模式的调用样例 examples/test_geir_segment_sum.cpp 中被完整复现:输入x为 shape{4, 2}、值全为 2.0 的 FLOAT32 Tensor,segment_ids为 shape{4}、值{0, 0, 1, 2}的 INT64 Const 节点,输出y的 shape 为{3, 2},与 README 中的示例一一对应。

关键语义要点

  • segment_ids必须按升序排序(但允许存在重复值,重复值表示多个行归入同一个段);
  • segment_ids指示当前分段的值归属于哪个段;
  • segment_ids的值必须>= 0,且各段编号可以不连续(例如[0, 0, 2]也是合法输入,此时y[1]将全为 0);
  • 输出 shape 为[max(segment_ids) + 1, x.shape[1:]],即输出第 0 维由最大段编号决定,其余维度与x保持一致。

参数说明

参数名输入/输出/属性描述数据类型数据格式
x输入输入数据,即公式中的xFLOAT32、FLOAT16、BFLOAT16、INT32、INT64、UINT32、UINT64ND
segment_ids输入分段索引,即公式中的segment_idsINT32、INT64ND
y输出输出值信息,即公式中的yFLOAT32、FLOAT16、BFLOAT16、INT32、INT64、UINT32、UINT64ND

从算子定义源码看,op_host/segment_sum_def.cpp 通过valueDataTypeXY声明了x/y支持的 7 种数值类型(FLOAT16、FLOAT、INT32、INT64、UINT32、UINT64、BF16),并通过valueDataTypeIdssegment_ids限定为 INT32、INT64 两种索引类型;输入输出均要求 ND 格式,并开启了动态 Shape、动态 Rank 支持。此外,segment_ids被标记为ValueDepend(OPTIONAL),表示索引张量的值在编译期可参与 tiling 决策。

约束说明

x

  • 维度至少为 1(rank >= 1);
  • 从算子定义看,x支持 1D~8D(见 op_graph/segment_sum_proto.h 注释),ND 格式下逐行累加。

segment_ids

  • 必须是 INT32 或 INT64 类型;
  • 必须为 1D Tensor,且segment_ids.shape[0] = x.shape[0](即每一行都有一个段归属);
  • 值必须按升序排序,且segment_ids.value >= 0

y

  • 类型必须与x相同;
  • 维度与x相同,shape 为[max(segment_ids) + 1, x.shape[1:]]

这些约束在 op_graph/segment_sum_proto.h 的@attention注释中逐条列出,同时单测与 ST 用例也对非法输入(如乱序索引)进行了覆盖。

调用说明

SegmentSum 当前支持的调用方式为GE 图模式:通过算子 IR 构图方式调用。

调用方式调用样例说明
GE 图模式test_geir_segment_sum.cpp通过算子 IR(segment_sum_proto.h)构图方式调用 SegmentSum 算子

算子 IR 定义

op_graph/segment_sum_proto.h 使用REG_OP宏注册算子原型:

REG_OP(SegmentSum) .INPUT(x, TensorType::NumberType()) .INPUT(segment_ids, TensorType::IndexNumberType()) .OUTPUT(y, TensorType::NumberType()) .OP_END_FACTORY_REG(SegmentSum)

xy使用NumberType(数值类型),segment_ids使用IndexNumberType(索引类型),与参数表中的类型范围一致。

GE 图模式调用示例详解

examples/test_geir_segment_sum.cpp 展示了完整的调用流程,核心步骤如下:

  1. 初始化 GE:通过ge::GEInitialize传入全局配置,包括ge.exec.deviceId=0ge.graphRunMode=1
std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}}; Status ret = ge::GEInitialize(global_options);
  1. 构造算子节点:创建op::SegmentSum("segmentSum")算子实例,并连接数据节点与常量节点:
auto segmentSum = op::SegmentSum("segmentSum"); // input x: Data 节点,shape {4, 2},填充值 2.0 vector<int64_t> xShape = {4, 2}; auto xData = op::Data("placeholder1").set_attr_index(0); TensorDesc xDesc(ge::Shape(xShape), FORMAT_ND, inDtype); xData.update_input_desc_x(xDesc); xData.update_output_desc_y(xDesc); graph.AddOp(xData); segmentSum.set_input_x(xData); // input segment_ids: Const 节点,shape {4},值 {0, 0, 1, 2} vector<int64_t> segmentIdsShape = {4}; auto segmentIdsConst = op::Const("placeholder2"); TensorDesc segmentIdsDesc(ge::Shape(segmentIdsShape), FORMAT_ND, DT_INT64); segmentIdsConst.SetAttr("value", segmentIdsTensor); graph.AddOp(segmentIdsConst); segmentSum.set_input_segment_ids(segmentIdsConst);
  1. 声明输出:为segmentSum设置输出y的 TensorDesc(shape{3, 2}),并加入输出算子列表;

  2. 构建并运行图:将图添加至ge::Session后调用RunGraph执行,随后读取输出 Tensor 数据、打印结果并落盘为.bin文件:

ret = session->AddGraph(graph_id, graph, graph_options); std::vector<ge::Tensor> output; ret = session->RunGraph(graph_id, input, output);

该样例还给出了数据生成辅助函数(如GenFloat32DataGenInt64Data)与GetDataTypeSize字节宽度映射,方便读者替换为任意受支持的数据类型组合进行验证。

源码实现机制

算子定义与注册(OpHost 侧)

op_host/segment_sum_def.cpp 中除了声明输入输出类型外,还通过OpAICoreConfig配置了算子行为:

OpAICoreConfig aicoreConfig; aicoreConfig.DynamicCompileStaticFlag(true) .DynamicFormatFlag(true) .DynamicRankSupportFlag(true) .DynamicShapeSupportFlag(true) .NeedCheckSupportFlag(false) .ExtendCfgInfo("opFile.value", "segment_sum_apt"); this->AICore().AddConfig("ascend950", aicoreConfig); this->AICore().AddConfig("ascend350", aicoreConfig);

其中DynamicShapeSupportFlag(true)DynamicRankSupportFlag(true)表示算子支持动态 Shape 与动态 Rank;ExtendCfgInfo("opFile.value", "segment_sum_apt")将 Kernel 实现指向 op_kernel/segment_sum_apt.cpp。

Tiling 策略(Host 侧)

Tiling 的入口在 op_host/arch35/segment_sum_tiling.cpp:

  • TilingPrepare4SegmentSum在编译期通过PlatformAscendC获取 AIV 核数(core_num)与 UB 内存大小(ub_size),存入SegmentSumCompileInfo
  • Tiling4SegmentSum调用TilingRegistry::GetInstance().DoTilingImpl(context)完成实际 tiling 计算。

op_host/arch35/segment_sum_tiling_base.h 中的SegmentSumBaseTiling定义了核心 tiling 参数,包括:外层维度outerDim_(即分段行数)、内层维度innerDim_(每行元素数)、段数量segmentNum_x数据类型字节数valueTypeBytes_、索引类型字节数idTypeBytes_等,为后续核内计算划分提供依据。

Kernel 实现(NPU 侧)

op_kernel/segment_sum_apt.cpp 是 SegmentSum 的 AICore 内核入口,通过 Tiling Key 分发到三种计算路径:

#define TEMPLATE_SIMT_TILING_KEY 1000 #define SIMD_ATOMIC_SUPPORT_TILING_KEY 2000 #define SIMD_DETERM_TILING_KEY 2002
  • TILING_KEY = 1000(SIMT 路径):调用SegmentSumSimt<DTYPE_X, DTYPE_SEGMENT_IDS>,适用于无原子操作的 SIMT 计算;
  • TILING_KEY = 2000(SIMD + 原子操作路径):先通过AllClear清空输出,再调用SegmentSumSimd执行求和。从代码看,该路径不适用于 UINT32、UINT64、INT64 三种类型(通过constexpr编译期分支排除,推测与原子加法类型支持范围有关);
  • TILING_KEY = 2002(SIMD 确定性路径):依次执行AllClear清空输出、SegmentSumSimdDeterm确定性求和、SegmentSumMultiCoreAdd多核结果累加,用于保证多核场景下求和结果的确定性与可复现性。

对应 Kernel 实现在 op_kernel/arch35/ 下:segment_sum_simt.hsegment_sum_simd.hsegment_sum_simd_determ.hsegment_sum_simd_mult_core_add.hclear_output.hsegment_sum_struct.h(tiling 结构体定义)。

算子二进制配置

op_host/config/ascend950/segment_sum_binary.json 列出了各 (x 类型, segment_ids 类型) 组合对应的二进制文件映射,共 14 组,覆盖 FLOAT32/FLOAT16/BFLOAT16/INT32/INT64/UINT32/UINT64 × INT32/INT64。所有条目的 shape 均为[-2](-2 表示动态 Shape),format 均为 ND,paramType 均为 required,说明该算子在 Ascend 950 上支持动态 Shape 的运行时编译。ascend350下另有同构的 segment_sum_binary.json。

测试与验证

SegmentSum 在仓库内提供了完整的测试资产与用例:

  • 输入生成脚本tests/assets/input.py:随机生成xsegment_ids。生成的索引值范围受输出第 0 维约束,随后调用np.sort保证segment_ids升序,符合算子约束要求;
  • Golden 脚本tests/assets/golden.py:直接调用 TensorFlow 的tf.math.segment_sum计算期望输出(bfloat16 输入先转为 float32 计算再转回),用于与 NPU 实际输出进行精度对拍;
  • ST 用例tests/st/arch35/ttk_kernel_segment_sum_st.csv:覆盖 FLOAT32×INT32/INT64、BFLOAT16×INT32、INT64×INT64 等多种类型组合,Shape 从 1D(如(5070,))到 8D(如(4, 3, 1, 7, 2, 5, 5, 5))不等,精度容忍度统一为 1e-8,同时验证了输出 shape 恒等于[max(segment_ids)+1, x.shape[1:]]
  • UT 用例tests/ut/op_host/arch35/test_segment_sum_tiling.cpp:针对 Tiling 逻辑做宿主侧单元测试。

小结

SegmentSum 是 ops-nn 中一个语义简洁但实现路径丰富的分段聚合算子:Host 侧通过动态 Tiling 感知 AIV 核数与 UB 容量,Kernel 侧依据数据规模在 SIMT 与多种 SIMD 路径间选择,并通过确定性累加保证多核结果可复现。使用时需严格遵守segment_ids升序、非负、1D 且长度等于x.shape[0]的约束,输出维度则由最大段编号决定。读者可基于 examples/test_geir_segment_sum.cpp 直接构造 GE 图进行功能验证,并借助仓库内的 tiling 单测与 ST 对拍用例深入理解其行为。

  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-nn
点击查看免费下载

相关推荐

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

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

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

立即咨询