已关闭
[Requirement|需求建议]: 【社区任务】Median算子开发交付(任务编号 04-9) #3513
bububu创建于  6月23日关闭于  7月29日
bububu
6月23日 创建

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

Backgroud(背景信息)

一、背景信息 (必填)

对标 PyTorch torch.median(昇腾官方对应接口为 aclnnMedian / aclnnMedianDim,由 Sort + Gather 等小算子拼接实现),在昇腾 NPU 上基于 Ascend C 编程语言实现功能一致的单算子融合版本,完成算子设计、开发、测试全流程工作,验收通过后将算子提交至昇腾算子开源仓 cann/ops-nn(合入路径 experimental/index/median)。

torch.median 有两种调用形态,需分别对齐:

  • 全局中位数 torch.median(input):将输入全部元素排序后取下中位数(元素数 n 为偶数时取排序后下标 (n-1)/2 处的值,与 NumPy 取均值语义不同),返回标量。
  • 按轴中位数 torch.median(input, dim, keepdim):沿 dim 轴取下中位数,返回 (values, indices),indices 为该中位数在原轴上的首个下标。

Origin(信息来源)

社区任务

Benefit / Necessity (价值/作用)

二、价值/作用 (必填)

官方 aclnn 拼接版需多次全量排序 + 中间张量落盘,HBM 往返多、kernel 启动多,是性能瓶颈点。本算子采用单算子融合的「选第 k 小」实现(不做全排序),一次搬入不落盘,显著降低 HBM 往返与跨核同步开销。
支持所有走入 aicore 的数据类型(fp16 / fp32 / bf16 / int16 / int32 / int64 / uint8 / int8)。必须实现算子泛化功能,满足各类合法输入场景的计算需求,验收阶段将采用泛化数据进行验收。

Design(设计方案)

三、设计方案 (必填)

host 侧设计

  1. 参数解析与校验

    • 输入:input,必选张量。
    • 属性:dim(int,MedianDim 用)、keepdim(bool)。
    • 输出:values(同 input dtype)、indices(int32,仅 MedianDim)。
    • 校验 input 非空;dtype ∈ {FLOAT16, FLOAT32, BF16, INT16, INT32, INT64, UINT8, INT8};dim 在 [-rank, rank) 合法范围内。
  2. tiling 策略

    • 经 L2 wrapper 把 reduce 轴换到末轴后,host 只需 reduce 长度 redLen、行数 batch = numel / redLen、下中位位 mid = (redLen-1)/2、dtype、UB 大小、AIV 核数等基础信息。
    • 行间相互独立,按 batch 占满 coreNum 核(放开旧版写死的 8 核);超大单行按 8192 对齐块分核。
    • workspace = 16MB 系统保留 + numel × 32(BigSort 有序 pair 余量)。
  3. tilingKey / 路径规划策略(按 dtype + redLen + batch 路由)

    • 全局 Median:BigSort 多核 K 路值划分(不排序、不整行进 UB),仅出 value,绝对早停。
    • MedianDim 多行、redLen ≤ 8192:MedianMain,Sort32 + 4-way MrgSort 批量排序取中,按 batch 满核。
    • 浮点大 L(redLen > 8192):L ≤ 24576 走 BigSel(整行常驻 UB + 向量化二分计数);batch == 1 走 BigSort 单行多核值二分;batch > 1 走 MedianBig 多行分块二分。
    • INT64:MedianHeap 分块标量值二分计数(避免整行 i64 进 UB 与 int16 索引溢出)。
    • dim 轴长度为 1 的退化场景:直通(values 即原值,index 恒 0)。

kernel 侧设计

  1. kernel 的实现流程

    • 分 Init 与 Process 两阶段,Process 含 CopyIn、Compute、CopyOut。
    • 从 GM 搬入 input tile 至 UB,按 dtype 选择计算路径:
      • float16 / float32 / int16 / int32:原精度比较。
      • bfloat16:Cast 到 float32 比较,结果再 Cast 回 bfloat16(Cast NONE+RINT 对齐 PyTorch)。
      • int8 / uint8:经 half / float 两步 Cast 中转后比较。
      • int64:标量分块值二分计数。
    • Compute 求第 k 小(k = mid 下中位数):Sort32 + MrgSort 排序取中 / 向量化值二分计数 / K 路值划分;需要 indices 时回扫取首个等值下标(对齐 torch)。
    • 多核归约:BigSort 每轮一次软件 SyncAll barrier,K = 8 路值划分把值域缩到 1/8,barrier 数约 20 降到约 6。
    • 将 UB 中的结果搬回 GM(values / indices)。
  2. AscendC 实现流程图

graph TD
    A["开始"] --> B["Host 校验 shape dtype dim 并计算 tiling"]
    B --> C["Kernel Init 读取 GM 地址与 tiling"]
    C --> D{"是否带 dim"}
    D -->|否 全局| E["BigSort 多核 K 路值划分 选第 k 小"]
    D -->|是 按轴| F{"redLen 是否大于 8192"}
    F -->|否| G["MedianMain Sort32 加 MrgSort 批量取中"]
    F -->|是| H["BigSel 或 MedianBig 向量化二分计数"]
    E --> I["CopyOut 写 values"]
    G --> J["回扫取首个等值下标 写 indices"]
    H --> J
    G --> I
    H --> I
    J --> I
    I --> K["结束"]

AscendC 实现与官方小算子拼接实现的差异及原因

  1. 官方 aclnnMedian / aclnnMedianDim 由 Sort + Gather 拼接、需全量排序并落盘中间张量;本算子单算子融合只「选第 k 小」,不全排序、一次搬入不落盘,减少 HBM 往返与 kernel 启动。
  2. 全局大 L 路径:官方多核值二分每轮一次 barrier、约 20 轮收敛;本算子 K = 8 路值划分把 barrier 降到约 6 轮。
  3. 多行按轴路径:本算子按 batch 占满核(行间独立);官方拼接受 Transpose / Sort 串行约束。
  4. 整型按轴:官方走全轴排序量级偏慢;本算子用值二分计数 / Sort32 取中,i8 / u8 / i16 / i64 普遍 10–1000× 反超。

支持硬件

产品 是否支持
Atlas A2 训练系列产品 / Atlas A3 系列产品 √

功能约束限制

参数名 含义 类型 数据类型 数据格式
input 输入张量 输入 FLOAT16、FLOAT32、BF16、INT16、INT32、INT64、UINT8、INT8 ND
dim 规约轴 属性 int -
keepdim 是否保留该轴 属性 bool -
values 中位数值 输出 同 input ND
indices 中位数在原轴的下标 输出 INT32 ND
  1. 全局偶数长度取下中位数(排序后下标 (n-1)/2),与 NumPy 取均值不同。
  2. indices 取首个等值下标,对齐 torch.median。
  3. dim 轴长度需支持 1(退化直通)。
likedislike
Bbububu
6月23日 关联了pull request:【CANN开源开放社区任务】【社区任务】AscendC实现Median算子贡献
oscillatedoscillated成员
6月23日 将 fullt 设为负责人
fulltower成员
7月7日 评论:

已安排审核,请关注PR检视意见

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