已关闭
[Requirement|需求建议]: 个人-AscendC实现Sort算子贡献 #288
zhoujianhua创建于  2025年12月18日关闭于  3月31日
zhoujianhua
zhoujianhua
2025年12月18日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

一、背景信息 (必填)

使用AscendC对TBE实现的Sort算子进行重构,实现了AscendC实现的Sort算子对Atlas 200/500 A2推理产品和Atlas 800I/T A2硬件的适配。

二、价值/作用 (必填)

Sort算子的主要功能是对指定张量沿目标维度执行升序/降序排序(未指定维度时默认沿最后一维排序),可返回排序后元素的原索引;对于多维张量(暂时只支持二维),可指定任意有效维度执行排序,排序后保持张量原有形状不变。在数学和工程领域中,张量排序是数据处理与数值计算的核心操作,它被广泛应用于深度学习(模型训练中的TopK筛选、注意力机制权重排序、损失函数异常值过滤、批量数据标准化前的排序预处理)、数值分析(大规模矩阵的行列排序、特征值排序、数值稳定性验证、行列式计算前的矩阵规整)等多个领域。Sort算子能够高效处理不同维度(1D/2D)、不同数据类型(float16/float32)的张量排序,支持整型维度参数输入、布尔型排序方向(升序/降序),适配昇腾硬件的并行计算特性(利用昇腾NPU多核并行加速,大张量排序性能显著优于CPU)。实现了Sort算子的AscendC实现,替代原有TBE算子在昇腾硬件上的适配,兼顾排序效率与内存利用率,支持自动微分(梯度通过排序后的值反向传播到原张量),满足昇腾平台下深度学习框架(MindSpore/CANN)的算子调用需求。

三、设计方案 (必填)

3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)

Aclnn直调

3.2 总体设计
3.2.1 算子支持的数据类型

float16、float两种数据类型。

3.2.2 host侧设计

Tiling分核策略:

  1. 获取形状与属性:
    • 获取输入张量的形状,得到高度 dimH 和宽度 dimW
    • 获取属性 axis(排序轴)和 descending(降序标志)。
  2. 确定并行任务维度:
    • axis == 0(按列排序),则独立排序任务数为 dimW(列数),单任务排序长度为 dimH
    • axis == 1(按行排序),则独立排序任务数为 dimH(行数),单任务排序长度为 dimW
  3. 计算可用核心数:
    • coreNum(硬件核心数)和独立任务数的较小值作为 workCoreNum
  4. 任务分配(大小核策略):
    • 将总任务数分配给 workCoreNum 个核心。
    • 计算 smallCoreNum(小核数量)和 bigCoreNum(大核数量)。
    • 计算 smallCoreDataNum(小核处理任务数)和 bigCoreDataNum(大核处理任务数,通常比小核多1)。
  5. 计算对齐与填充参数:
    • sliceLen:实际需要排序的元素个数。
    • realSortLen:满足 Sort 指令要求的 32 字节对齐长度(((sliceLen + 31) / 32) * 32)。
    • align8:8 字节对齐长度。
    • padLenalign8 - sliceLen,用于 DataCopyPad 的填充长度。
    • dupCount32 - align8 % 32,用于尾部额外填充的元素个数。
  6. 结构体写回:
    • 将上述计算出的核心分配参数、维度信息、填充参数写入 SortV2TilingData 结构体。
3.2.3 kernel侧设计

进行 Init 和 Process 两个阶段,其中 Process 包括数据搬入(CopyIn)、计算(Compute)、数据搬出(CopyOut)三个阶段。

初始化(Init)

  1. 核类型判断与任务分配:
    • 获取当前核索引 blockIdx = AscendC::GetBlockIdx()
    • 根据 Tiling 数据判断当前核是大核还是小核,计算当前核负责的任务起始索引 startSlice 和任务数量 coreDataNum
    • 计算全局内存偏移 coreDataStart
      • axis == 0,偏移为 startSlice(列偏移)。
      • axis == 1,偏移为 startSlice * dimW(行偏移)。
  2. GM Tensor 映射:
    • 设置 xGmindexGmyGmdstIndexGm 的全局缓冲区地址,指向当前核负责的数据段。
  3. 缓冲区初始化:
    • 根据 realSortLen 计算所需的 buffer 大小。
    • 初始化队列:inQueueX(输入数据与索引)、outQueueY(输出数据)、dstIndexQ(输出索引)、calcQ/tmpQ/concatQ(计算临时空间)。

数据搬入(CopyIn)

  1. 按轴处理:
    • Axis=0(列排序): 数据在 GM 中不连续(跨行)。循环 dimH 次,每次使用 DataCopyPad 搬运一个元素到 xLocalindexLocal 的对应位置。
    • Axis=1(行排序): 数据在 GM 中连续。使用 DataCopyPad 一次性搬运一行数据。
  2. 填充处理(Padding):
    • 利用 DataCopyPadpadParams 进行填充。
    • 填充值根据 descending 属性决定:若降序则填充 -FLT_MAX,若升序则填充 FLT_MAX,确保填充值排在最后。
    • 若有剩余对齐需求(dupCount > 0),使用 Duplicate 指令进行额外填充。

计算流程(Compute)

  1. 数据预处理:
    • inQueueX 取出数据。
    • 使用 Concat 指令将数据按照 Sort 指令要求的格式进行拼接(concatRepeat)。
    • 升序转换: 若为升序排序(!descending),调用 Muls 乘以 -1,将数值取反,从而利用硬件的降序排序能力实现升序。
  2. 排序执行:
    • 调用 AscendC::Sort 指令对数据和索引进行全排序。
  3. 后处理:
    • 调用 AscendC::Extract 提取排序后的有效结果。
    • 数值恢复: 若进行了升序转换,再次调用 Muls 乘以 -1 恢复原始数值。
  4. 结果入队:
    • 将排序后的值放入 outQueueY,索引放入 dstIndexQ

数据搬出(CopyOut)

  1. 按轴回写:
    • Axis=0(列排序): 循环 dimH 次,使用 DataCopyPadyLocaldstIndexLocal 中的元素逐个写回全局内存的对应跨行位置。
    • Axis=1(行排序): 使用 DataCopyPad 将排序后的一行连续写回全局内存。
3.3 支持硬件

Atlas 200/500 A2推理产品和Atlas 800I/T A2硬件

3.4 算子约束限制
  • 支持 float16, float32 数据类型。
  • 输入 shape 暂时只支持二维,可以指定其中任一维度排序,输出排序结果以及排序后的索引顺序(可选)。
  • 支持升序和降序排序,排序的稳定性取决于sort接口。
  • 不支持广播机制,仅支持对给定维度的独立排序。
  • 在 UB 为 192KB 的情况下,预估支持的最大排序元素个数约为 1500~2000 个(基于 float 类型估算,具体取决于 BUFFER_SIZE 定义及数据类型)。超出此长度的 Shape 当前版本暂不支持。

💡 备注(选填)

likedislike
zhoujianhua
zhoujianhua
2025年12月18日 评论:

/assign

likedislike
CANN-robotCANN-robot成员
2025年12月18日 将 LePenseur 设为负责人
zhoujianhuazhoujianhua
2025年12月18日 关联了pull request:个人-AscendC实现Sort算子贡献
zhoujianhuazhoujianhua
1月27日 修改了issue 的描述
CANN-robot
CANN-robot成员
3月18日 评论:

Notice

This issue is already assigned to LePenseur. Please do not assign repeatedly.

likedislike
CANN-robot
CANN-robot成员
3月19日 评论:

Notice

This issue is already assigned to LePenseur. Please do not assign repeatedly.

likedislike
CANN-robotCANN-robot成员
3月31日 关闭了 issue
CANN-robotCANN-robot成员
3月31日 添加了label:resolved