已关闭
[Requirement|需求建议]: TopkV2算子在小k情况下支持双调排序的实现 #2898
cy_hw创建于 12 天前关闭于 7 天前
12 天前 添加了label:requirement
12 天前 关联了pull request:TopkV2算子支持小k场景下的双调排序
陈思
11 天前 评论:
11 天前 评论:
/assign @caoyan_huawei


11 天前 将 caoyan_huawei 设为负责人
10 天前 修改了issue 的描述
9 天前 修改了issue 的描述
9 天前 修改了issue 的描述
7 天前 关闭了 issue
7 天前 添加了label:resolved
TopKV2 支持双调排序(Bitonic Sort)需求设计
1. 背景与目标
TopKV2 算子在
sorted=True场景下需要对 TopK 候选结果进行最终排序。传统实现使用AscendC::Sort高阶 API(基于 Radix Sort),其排序网络规模与数据量正相关,在小 K 值(2 ≤ k ≤ 32)场景下存在固定开销占比过高的问题。本特性引入固定规模为 32 的双调排序网络(Bitonic Sorting Network),利用 arch35 芯片的 SIMD 向量寄存器(Reg)和 SIMT warp 级通信能力,在小 K 场景下以更低的延迟完成收尾排序。
2. 触发条件
三个条件同时满足时激活双调排序路径:
sorted = TrueisSort2 ≤ k ≤ 32kValue判定sortPolicy = 1Host 侧由
IsBitonicSmallTopkMode(kValue, sortPolicy, isSort)统一判定。3. 系统架构
3.1 总体数据流
3.2 Tiling Key 机制
Host 侧通过偏移量叠加生成 tiling key,Kernel 侧通过
TILING_KEY_IS宏匹配:MERGE_SORT_TILING_OFFSETBITONIC_SORT_TILING_OFFSET两个偏移可叠加(如 FLOAT + Bitonic + Merge = 3003 + 50000 + 10000 = 63003),支持"小轴 + 小 K + 排序"复合场景。
完整 key 映射表:
3.3 IS_BITONIC_SORT 模板参数
Kernel 侧通过编译时模板参数
IS_BITONIC_SORT(bool)控制双调路径的启用:// top_k_v2_apt.cpp 入口 #if TILING_KEY_VAR == TOPK_COMMON_TILING_KEY_FLOAT generateOpObject<float, uint32_t, topkV2::B32_BITE_SIZE, DTYPE_INDICES, false>(...); #elif TILING_KEY_VAR == TOPK_BITONIC_TILING_KEY_FLOAT generateOpObject<float, uint32_t, topkV2::B32_BITE_SIZE, DTYPE_INDICES, true>(...); #endifIS_BITONIC_SORT传递到 7 类算子模板,覆盖 TopKV2 全部计算路径:RadixSortTopKradix_sort_top_k.hSortTopKRes后调用RunBitonicSmallTopKFinalizeRadixSortTopKSingleBlockradix_sort_top_k_single_block.hProcessSingleTime后调用RunBitonicSmallTopKFinalizeRadixSortTopKSingleCoreradix_sort_top_k_single_core.hSortTopKRes后调用RunBitonicSmallTopKFinalize;IsUseRadixGather编译时短路RadixSortTopKMultiCoreOptimizationradix_sort_top_k_inter_core_template_optimization.hProcessSingleLoop后调用RunBitonicSmallTopKFinalizeSortAndTopKMoreCoresort_and_top_k_more_core.hGetBitonicSmallTopKRes独立路径替代GetTopKResMergeSorttop_k_merge_sort.hVbsMergeSort后调用RunBitonicSmallMergeSortFinalizeNonLastTopKNonLastSmallAxisNonTransposetop_k_non_last_small_non_transpose.hRunTopK/RunMergeSort后调用RunBitonicSmallTopKFinalizeNonLast3.4 编译时优化
原实现中
IsUseRadixGather通过运行时成员变量sortPolicy_判断是否启用 Gather 路径。重构后:IS_BITONIC_SORT = true时,IsUseRadixGather在if constexpr中编译期短路返回IsBitonicSmallTopkMode已保证IS_BITONIC_SORT = true蕴含IS_SORT、k ∈ [2, 32]、sortPolicy == 1,逻辑等价sortPolicy_成员变量及其 tiling data 读取// top_k_util_type_simd.h template <bool IS_SORT, bool IS_BITONIC_SORT> __aicore__ inline bool IsUseRadixGather(bool isB32OrB64Integer, bool needSortWithIndex) { if constexpr (!IS_SORT || !IS_BITONIC_SORT) { return false; // 编译时短路 } return isB32OrB64Integer && !needSortWithIndex; }4. 核心实现:top_k_small_size_bitonic_sort.h
4.1 文件整体架构
4.2 两条实现路线
__simd_callee__Reg::Gather,Reg::Select,Reg::Compare__simt_callee__/__simt_vf__asc_shfl_xor,asc_ballot,__popc1 字节类型不能走 Reg 路径的原因:arch35 的 SIMD
Reg::API(Gather/Compare/Select/LoadAlign 等)最小操作位宽为 16 位,不支持 8 位RegTensor。1 字节类型无法加载进 Reg 寄存器体系,只能回退到 SIMT 标量线程路径,用asc_shfl_xor替代Reg::Gather。这是硬件 ISA 的功能限制,非性能选择。4.3 值到 Key 的映射(BitonicSmallRegBuildKey32)
双调网络用整数比较替代浮点/有符号数比较,需将各种类型的数据位模式映射为可单调比较的
uint32key:key = rawBitskey = rawBits XOR signBitkey = rawBits XOR signBitkey = rawBits XOR allBitskey XOR allBitsNaN 特殊处理:NaN 的指数全 1 且尾数非 0。
IsLargest时 NaN key 设为allBits(排在最后),否则设为 1(排在最前),保证 NaN 总在结果末尾。浮点 -0 会被清零,与 +0 统一处理。4.4 双调排序网络
32-lane 双调网络由 5 级归并构成,每级内部按递减 stride 做 compare-swap,共 15 个 SwapStage:
先构建双调序列(size 递增),再做双调合并(stride 递减),最终全排序。
4.4.1 Reg 路径 Compare-Swap(BitonicSmallReg32SwapStage)
每个 lane 执行以下步骤:
peerLane = lane XOR strideReg::Gather获取对端 value/key/index/groupCompareGroup=true时按 (group, index) 字典序;false时按 key 大小swapMask = (compare && lowValid) || !highValidSize==32时全升序;否则按子网方向takePeerMask = !(swapMask XOR directionMask)4.4.2 SIMT 路径 Compare-Swap(BitonicSmallSwapStage)
每个 lane 线程持有一个标量元素:
asc_shfl_xor获取对端 lane 的 value/index4.5 候选选择与收尾排序
4.5.1 已排序候选收尾(BitonicSmallFinalizeSelection)
处理 TopK/RadixSelect 输出的已按值排序的候选集:
4.5.2 未排序候选收尾(BitonicSmallFinalizeExactSelection)
处理 Gather 路径输出的未排序精确候选集。与上述区别在于候选集未按值排序,需要先归约求阈值:
4.5.3 Reg 路径收尾(BitonicSmallRegFinalizeSelection)
Reg 路径的收尾逻辑与 SIMT 路径语义等价,但用向量寄存器指令实现:
UINT32_MAX标记无效BitonicSmallRegBuildKey32)BitonicNetwork<true>),再按 key 全排序(BitonicNetwork<false>)按
sizeof(T)分发到不同特化路径:BitonicSmallRegFinalizeSelectionB64BitonicSmallRegFinalizeSelectionB164.6 阈值 Gather 收集(BitonicGatherThresholdTileKernel)
针对大轴场景,使用 SIMT warp 内 ballot + popcount 做前缀和,按阈值收集候选元素:
4.7 多行批量调度
SIMT 路径支持多行并行处理,
threadIdx.y为行索引,threadIdx.x为 lane (0~31):RunBitonicFinalizeSelectionRowsBITONIC_SMALL_TOPK_MAX_ROWS(32)RunBitonicFinalizeSmallSourceRowsBITONIC_SMALL_TOPK_MAX_ROWS(32)分批调度:
for (rowStart = 0; rowStart < rowCount; rowStart += 32)5. 模板集成点详解
5.1 尾轴场景
RadixSortTopK / RadixSortTopKSingleCore / RadixSortTopKSingleBlock
这三类模板在
AscendC::Sort或AscendC::TopK高阶 API 完成初步排序后,通过if constexpr (IS_BITONIC_SORT)调用收尾排序:// radix_sort_top_k.h:892 AscendC::Sort<T, T_INDEX_TO, false, sortConfig>(topkValueOutLocal, topkValueOutIndexLocal, ...); if constexpr (IS_BITONIC_SORT) { RunBitonicSmallTopKFinalize<T, T_INDEX_TO, IS_LARGEST>( topkValueOutLocal, topkValueOutIndexLocal, topkValueInitInput_, 1U, ...); }调用时机:高阶 API 输出 k 个候选值和索引 → Bitonic 收尾排序 → 写回 GM。
RadixSortTopKSingleCore 的 Gather 优化
当
IS_BITONIC_SORT = true且数据类型为 int32/uint32/int64/uint64 时,IsUseRadixGather编译时返回true,走 Gather 路径(GatherTileTopK2CopyOut)替代传统 TopK API,再用RunBitonicFinalizeExactSelection收尾:// radix_sort_top_k_single_core.h:394-398 if constexpr (IS_SORT && std::is_integral_v<T>) { if (useRadixGather_) { GatherTileTopK2CopyOut(loopTime, involvedDataMask); return; } }RadixSortTopKMultiCoreOptimization
// radix_sort_top_k_inter_core_template_optimization.h:232 AscendC::TopK<T, true, false, false, TopKMode::TOPK_NORMAL, topkConfig>(...); if constexpr (IS_BITONIC_SORT) { RunBitonicSmallTopKFinalize<T, int32_t, IS_LARGEST>(topkOutValue, topkOutIndexValue, ...); }SortAndTopKMoreCore
GetBitonicSmallTopKRes是独立路径,完全替代GetTopKRes:// sort_and_top_k_more_core.h:295 if constexpr (IS_BITONIC_SORT) { GetBitonicSmallTopKRes(sortLoopRound); return; }该路径逐行处理:MTE2 加载 → V 计算(
RunBitonicSmallTopKFinalize)→ MTE3 输出,行间通过MTE3_MTE2事件同步防止 UB 覆盖。5.2 非尾轴场景
TopKNonLastSmallAxisNonTranspose
非尾轴场景使用
AscendC::TopK或AscendC::Sort高阶 API 后,调用专用收尾入口:// top_k_non_last_small_non_transpose.h:259 RunBitonicSmallTopKFinalizeNonLast<SortT, IsLargest, IS_BITONIC_SORT>( topkValue_, topkIndex_, this->sortInput_, this->axisLen_, kValue_, ...);该入口的分发逻辑:
RunBitonicFinalizeSmallSourceRows(小源行)RunBitonicFinalizeSelectionRows(SIMT 行批量)BitonicSmallRegFinalizeSelection(Reg 逐行)MergeSort
// top_k_merge_sort.h:274 RunBitonicSmallMergeSortFinalizeNonLast<SortT, IsLargest, IS_BITONIC_SORT>(dst, dstIndex, ...);分发逻辑:
RunBitonicFinalizeSelectionRows(SIMT)BitonicSmallRegFinalizeSelection(Reg)6. 数据类型支持矩阵
7. 流水线同步设计
7.1 SortAndTopKMoreCore 逐行处理同步
GetBitonicSmallTopKRes逐行处理时,valuesLocal来自depth=1队列,indicesLocal来自固定TBuf。行间需防止 MTE3 读 UB 与下一行 MTE2 写 UB 的竞争:// sort_and_top_k_more_core.h:282-287 DataCopyPad(this->outValueGm_[outputOffset], valuesLocal, valueCopyParams); DataCopyPad(outputIndex[outputOffset], indicesLocal, indexCopyParams); event_t mte3ToMte2 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2)); SetFlag<HardEvent::MTE3_MTE2>(mte3ToMte2); WaitFlag<HardEvent::MTE3_MTE2>(mte3ToMte2); topkResQueue_.FreeTensor(valuesLocal);V_MTE3只保证 V 先于 MTE3 启动,不保证 MTE3 完成后下一轮 MTE2 才写 UB。MTE3_MTE2同步确保上一行 MTE3 CopyOut 完成后,下一行 MTE2 CopyIn 才开始写 UB。末尾SyncAll()仅同步核间,无法补救核内行间竞争。7.2 Reg 路径内存屏障
Reg 路径在加载前和存储前使用
Reg::LocalMemBar确保向量寄存器的加载/存储有序:// top_k_small_size_bitonic_sort.h:383 Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); // 加载、计算、存储...8. 关键设计决策
8.1 固定规模 32 的双调网络
双调网络规模固定为 32(
BITONIC_SMALL_TOPK_SIZE),与 arch35 的 warp 大小对齐。k < 32 时通过 pad + 无效标记(UINT32_MAX)处理,无效元素在排序中自然被挤到末尾。这简化了实现并充分利用硬件并行度。8.2 值到 Key 映射
用整数比较替代浮点/有符号数比较,使得双调网络中统一用
uint32比较指令完成排序。浮点 NaN 通过特殊 key 值处理,保证排序结果的一致性和可预测性。8.3 两轮排序策略
收尾排序采用两轮双调网络:
CompareGroup=true):按 (group, index) 排序,恢复原始 lane 顺序CompareGroup=false):按值排序,输出最终有序结果这是为了兼容 TopK 高阶 API 的输出语义——当存在重复值时,TopK 选择的候选可能与 Bitonic gather 的候选不同,需要先恢复原始顺序再重新排序。
8.4 IS_BITONIC_SORT 编译时参数
选择编译时模板参数而非运行时变量,使得:
if constexpr在编译期消除非 Bitonic 路径的代码IsUseRadixGather等条件判断编译期短路,零运行时开销sortPolicy字段的读取和存储9. 文件依赖关系