已关闭
[Requirement|需求建议]: 910b的sort_with_index实现 #1765
zhaohujie创建于  6月2日关闭于  6月23日
zhaohujie
zhaohujie
6月2日 创建

一、背景信息 (必填)

仓库已存在完整算子 math/sort_with_index,但其 AICore().AddConfig 仅注册 ascend950,kernel 全在 op_kernel/arch35/(ascend950 专用),在 Ascend910B 上无可用 kernel/二进制(系统 opp 中的 SortWithIndex 同样仅 arch35,且未导出 aclnnSortWithIndex)。本需求补齐 Ascend910B 原生 AscendC 开源实现,落到 experimental/math/sort_with_index/,以 math/sort_with_index 为接口真值源。

二、价值/作用 (必填)

  • SortWithIndex 在 Atlas A2(910B)上可用,填补该芯片上的能力缺口。
  • 语义:沿指定维度排序,双输出——排序后的 values(y)与跟随排序重排后的索引(sorted_index)。等价于「排序 + 携带一个调用方提供的索引张量同步搬移」,是 Top-K、argsort、检索/重排、稀疏索引跟随等场景的基础算子(index=0..N-1sorted_indextorch.sort 的 indices)。
  • 支持 descending(降序)与 stable(稳定排序)语义。

三、设计方案 (必填)

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

Aclnn 直调(aclnnSortWithIndex 两段式)、GE 图模式(proto,动态 shape/rank)、PyTorch(经 torch_npu)。

3.2 总体设计

基于 910B 可用的 Concat/Sort/Extract(大轴核内 MrgSort)重写排序核;索引跟随利用 AscendC Sort 的位置稳定性。

3.2.1 算子支持的数据类型
输入/输出 value(x、y) 输入/输出 index(index、sorted_index)
float16 / float32 / bfloat16 / int32 int32
3.2.2 host侧设计
  • op_defAddConfig("ascend910b", …),声明 4 组 dtype,属性 axis(-1)/descending(false)/stable(false)。
  • infershapey.shape = x.shapesorted_index.shape = index.shape;静态 shape 下强制校验 x.shape == index.shape
  • tiling:按「非排序维(行)」多核切分、行内不跨核;UB 切分 + dtype 感知 padLen 对齐;TilingKey 三维 (VALUE_DT, INDEX_DT, SIZE_MODE)SIZE_MODE ∈ {单 tile, 大轴 MrgSort, 空 tensor};单行排序轴长超出核内上限时 tiling 返回失败,不崩溃。
  • op_graph:proto + InferDataType(输出 dtype 跟随输入)。
3.2.3 kernel侧设计
  • 排序核:Concat + Sort + Extract(基于 8B proposal 记录),大轴用核内 MrgSort 多块归并(无 SyncAll)。
  • 索引跟随:AscendC Sort 在分数相等时按输入位置(positional stable)定序,而非按索引通道值;故 int32 调用方索引可直接进入 Sort 索引通道并提取得到 sorted_index,省去 position 通道 + 偏移 Muls + Gather(int64 与大轴 MrgSort 路径仍用 position 通道 + Gather)。
  • dtype 通路:half/float 直接 Sort;bf16/int32-value 经 Cast → float → Sort → Cast 回原 dtype。
  • 升序经 Muls(-1) 取反;±Inf 用 ±Inf 哨兵(升序 +Inf / 降序 -Inf,避免有效 ±Inf 被哨兵顶替)。
  • 同步精简:尾部屏障按真实依赖收窄(PIPE_V,仅覆盖 V→MTE3)。
3.3 支持硬件

Ascend910B(Atlas A2 系列,DAV_2201)

3.4 算子约束限制

  • index 仅 int32:910B 框架排序算子族在非 RegBase 上强制 index 输出为 int32,int64-index 经 aclnn 不可选中,故首版不对外暴露 int64-index(int64 kernel 代码保留但不注册)。
  • 仅最后一维排序axis 须为 -1rank-1,否则报错;行间独立。
  • 单行排序轴长有上限(dtype 相关,约 2500–3000,随运行时可用 UB 取值):超界返回 tiling 失败(优雅拒绝,非崩溃);超大轴的 GM workspace 方案留后续迭代。
  • int32 value 仅 |x| ≤ 2^24 精确(经 float 路径排序)。
  • NaN 升序落序列开头(910B 实现,区别于 torch.sort 的末尾约定);位型比对按 isnan(非按 bit)。±Inf:+Inf 升序末尾、-Inf 升序开头。
  • 不支持广播xindex 必须同 shape。

💡 备注(选填)

likedislike
zhaohujie
zhaohujie
6月2日 评论:

/assign @zhaohujie

likedislike
CANN-robotCANN-robot成员
6月2日 将 zhaohujie 设为负责人
zhaohujiezhaohujie
6月2日 修改标题为 “[Requirement|需求建议]: 910b的sort_with_index实现”,原标题为“[Requirement|需求建议]: 910b的sortwithindex实现”
zhaohujiezhaohujie
6月2日 修改了issue 的描述
zhaohujiezhaohujie
6月2日 关联了pull request:feat(sort_with_index): 910b的sort_with_index实现
CANN-robotCANN-robot成员
6月23日 关闭了 issue
CANN-robotCANN-robot成员
6月23日 添加了label:resolved