已关闭
[Requirement|需求建议]: 新增 Median/NanMedian 算子并完善 KthValue/Sort 特殊值排序语义 #2968
黄晓彬创建于  10 天前关闭于  8 天前
黄晓彬成员
10 天前 创建

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

一、背景信息 (必填)

当前 ops-math 在 Ascend 950 上已有 KthValue、Sort、TopKV2 等有序选择/排序基础能力,但缺少沿指定维度同时返回中位数值和源索引的 Median、NanMedian 算子。直接为两个新算子各自维护完整排序实现会重复调度、workspace 和稳定索引逻辑,也容易造成 NaN、重复值及 ±0 等特殊值语义不一致。

本需求新增 Median 和 NanMedian,并复用 KthValue 的多调度选择框架:

  • Median 取 lower median;浮点输入任一约简行包含 NaN 时返回该行 NaN。
  • NanMedian 忽略 NaN 后取非 NaN 元素的 lower median;全 NaN 行返回 NaN。
  • 两个算子均返回 values 和源轴 indices,约简维保留为 1。
  • 同步完善共享 Sort/KthValue 基础逻辑中 NaN key、重复值及 ±0 的稳定顺序,避免 Median、NanMedian、KthValue、Sort、TopKV2 的特殊值行为分叉。

二、价值/作用 (必填)

  1. 补齐 Ascend 950 上 Median/NanMedian 的 AICore 原生实现,支持动态 shape、动态 rank、负 dim 和非末轴场景。
  2. 复用 KthValue 的 insertion、two-stage、merge、radix-select、非末轴 small-axis 等调度,避免为中位数再执行完整排序,并减少重复代码。
  3. 对齐 PyTorch median/nanmedian 的 lower-median、NaN 和稳定索引语义,覆盖重复值、signed zero、普通 NaN 及全 NaN 数据。
  4. 整理 KthValue kernel 统一 dispatch,并补充 KthValue、Sort、TopKV2 的 golden 与 TTK 回归数据,降低共享排序基础设施修改带来的回归风险。

三、设计方案 (必填)

3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)
  • 新增 Median/NanMedian GE OpDef、InferShape、Tiling、AICore kernel 入口和 L0 op API。
  • 算子模块依赖 kth_value 与 sort,共享 arch35 kernel/tiling 基础设施。
  • 支持动态编译静态化、动态 shape 和动态 rank,输入输出格式为 ND。
  • 当前 CMake 使用 aclnn_exclude,本需求不新增独立的公开 aclnn wrapper;由仓内 L0/GE 注册通路使能。
3.2 总体设计

以 KthValue 为统一选择引擎,在 tiling data 中增加 medianMode

  • STATIC:原 KthValue/整数 Median/NanMedian 的固定 k 路径。
  • PROPAGATE_NAN:浮点 Median;行内存在 NaN 时选择稳定排序后的首个 NaN。
  • IGNORE_NAN:浮点 NanMedian;按非 NaN 数量动态计算 lower-median rank。

Host 侧继续选择既有 KthValue schedule;Kernel 侧在各 schedule 的加载/选择阶段统一完成 NaN 归一化、有效元素计数和实际 rank 解析。整数类型不进行 NaN 扫描,保持原静态 k 快路径。

3.2.1 算子支持的数据类型

Median/NanMedian 输入 x 与输出 y 支持:

  • float16、float32、bfloat16
  • int8、uint8、int16、uint16、int32、uint32、int64、uint64

indices 固定为 int64;所有输入输出均为 ND。属性 dim 为可选 int64,默认 -1。

3.2.2 host侧设计
  1. OpDef/InferShape

    • 校验输入 rank、dim 和约简轴长度。
    • 规范化负 dim;values/indices 输出保持输入 rank,并把约简维设置为 1。
    • 注册 Ascend 950 的 AICore kernel、动态 shape/rank 能力及 simplified key。
  2. Tiling

    • lower median 使用零基 rank (axisLen - 1) / 2;内部沿用 KthValue 的一基 k 参数。
    • 根据输入 dtype 和算子类型设置 medianMode:浮点 Median 传播 NaN,浮点 NanMedian 忽略 NaN,整数统一走 STATIC。
    • 复用 KthValue 的末轴/非末轴、单核/多核、small-axis/merge/radix 路由及原有 UB 规划。
    • multi-core merge/radix 的动态 median rank 需要跨核汇总时,在既有算法 workspace 后追加按 block 对齐的 per-core uint32 计数区;不改变整数静态 k 路径的 workspace。
  3. 公共化

    • 将 KthValue 原 kernel 入口中的 dtype/schedule 分发抽取为统一 dispatch,Median、NanMedian 仅提供独立 kernel 入口和 tiling key,避免复制算法主体。
3.2.3 kernel侧设计
  1. NaN 与 rank

    • 浮点数据进入 median 模式后把不同 payload/sign 的 NaN 归一为 canonical quiet NaN key,使 NaN 形成确定后缀。
    • Median 检测行内 NaN:存在时选首个稳定 NaN;不存在时使用静态 lower-median rank。
    • NanMedian 统计非 NaN 数量并计算 (nonNanCount - 1) / 2;全 NaN 行保留 NaN 结果。
    • 对已排序 UB 行使用二分定位 NaN 后缀起点;对分块/多核路径使用 UB 标量或 workspace 聚合计数。
  2. 调度复用

    • insertion、small-axis two-stage、merge intra-core、merge more-core、radix one-core/more-core、radix-select 和 non-last small-axis 均通过统一 helper 解析实际 rank。
    • 整数类型在编译期裁剪 NaN 处理,避免引入额外扫描和性能开销。
  3. 稳定顺序

    • Sort 公共基类补充 signed-zero source-order 处理,统一 ±0 的稳定索引选择。
    • FP16/BF16/FP32 radix key 对 NaN 使用 canonical key;保持原数据值不被无关改写。
    • Golden 对重复值/等价 key 的 indices 使用稳定排序语义;TopKV2 对并列值允许按合法稳定集合比对。
3.3 支持硬件
  • Ascend 950(arch35)
  • AICore Mix AIV kernel
  • ND 格式

3.4 算子约束限制

  • 输入 rank 必须大于 0,dim 范围为 [-rank, rank - 1]
  • 约简轴不能为空。
  • 非末轴当前支持轴长 [2, 2048];末轴沿用 KthValue 已有调度能力。
  • 输出保留约简维,长度固定为 1。
  • Median/NanMedian 均采用 lower median,不对偶数长度的中间两个值求平均。
  • 当前不新增独立公开 aclnn 接口和接口 README,交付范围为 OpDef、L0、Tiling、Kernel、InferShape、golden、UT/TTK 用例。

💡 备注(选填)
验收重点覆盖 11 种 dtype、末轴/非末轴、动态 shape/rank、负 dim、重复值、signed zero、普通 NaN、全 NaN,以及 KthValue/Sort/TopKV2 共享路径回归。

likedislike
黄晓彬成员
10 天前 添加了label:requirement
CANN-robotCANN-robot成员
8 天前 关闭了 issue
CANN-robotCANN-robot成员
8 天前 添加了label:resolved