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["结束"]
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 侧设计
参数解析与校验
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)合法范围内。tiling 策略
redLen、行数batch = numel / redLen、下中位位mid = (redLen-1)/2、dtype、UB 大小、AIV 核数等基础信息。batch占满coreNum核(放开旧版写死的 8 核);超大单行按 8192 对齐块分核。workspace = 16MB 系统保留 + numel × 32(BigSort 有序 pair 余量)。tilingKey / 路径规划策略(按 dtype + redLen + batch 路由)
L ≤ 24576走 BigSel(整行常驻 UB + 向量化二分计数);batch == 1走 BigSort 单行多核值二分;batch > 1走 MedianBig 多行分块二分。kernel 侧设计
kernel 的实现流程
Init与Process两阶段,Process含 CopyIn、Compute、CopyOut。inputtile 至 UB,按 dtype 选择计算路径:float16 / float32 / int16 / int32:原精度比较。bfloat16:Cast 到float32比较,结果再 Cast 回bfloat16(Cast NONE+RINT 对齐 PyTorch)。int8 / uint8:经half / float两步 Cast 中转后比较。int64:标量分块值二分计数。indices时回扫取首个等值下标(对齐 torch)。SyncAllbarrier,K = 8 路值划分把值域缩到 1/8,barrier 数约 20 降到约 6。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 实现与官方小算子拼接实现的差异及原因
aclnnMedian / aclnnMedianDim由 Sort + Gather 拼接、需全量排序并落盘中间张量;本算子单算子融合只「选第 k 小」,不全排序、一次搬入不落盘,减少 HBM 往返与 kernel 启动。batch占满核(行间独立);官方拼接受 Transpose / Sort 串行约束。支持硬件
功能约束限制
(n-1)/2),与 NumPy 取均值不同。indices取首个等值下标,对齐torch.median。dim轴长度需支持 1(退化直通)。