已关闭
[Requirement|需求建议]: TopkV2算子在小k情况下支持双调排序的实现 #2898
cy_hw创建于  12 天前关闭于  7 天前
cy_hw
cy_hw成员
12 天前 创建

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 = True 算子属性 isSort 要求输出有序
2 ≤ k ≤ 32 Host 侧 kValue 判定 双调网络固定 32-lane,k 超出则不适用
sortPolicy = 1 算子属性 显式启用 Bitonic Small TopK 路径,默认 0

Host 侧由 IsBitonicSmallTopkMode(kValue, sortPolicy, isSort) 统一判定。

3. 系统架构

3.1 总体数据流

Host 侧 Tiling                          Kernel 侧 Dispatch
┌─────────────────────────┐             ┌────────────────────────────────┐
│  IsBitonicSmallTopkMode  │             │  TILING_KEY_IS(COMMON_KEY)     │
│          ↓               │             │  TILING_KEY_IS(BITONIC_KEY)     │
│  dataTypeKey += 50000    │ ──────────→ │          ↓                     │
│  (BITONIC_SORT_OFFSET)   │  tiling key │  IS_BITONIC_SORT=false / true │
│          ↓               │             │          ↓                     │
│  SetTilingKey(key)       │             │  generateOpObject<...,IS_..>   │
└─────────────────────────┘             └────────────────────────────────┘
                                                    │
                                     ┌──────────────┼──────────────┐
                                     ↓              ↓              ↓
                               RadixSortTopK   MergeSort    SortAndTopK
                                     │              │              │
                                     ↓              ↓              ↓
                            RunBitonicSmallTopKFinalize (核心入口)

3.2 Tiling Key 机制

Host 侧通过偏移量叠加生成 tiling key,Kernel 侧通过 TILING_KEY_IS 宏匹配:

偏移类型 说明
基础 key 1001~4002 由数据类型决定
MERGE_SORT_TILING_OFFSET +10000 Merge Sort 模式
BITONIC_SORT_TILING_OFFSET +50000 Bitonic 模式

两个偏移可叠加(如 FLOAT + Bitonic + Merge = 3003 + 50000 + 10000 = 63003),支持"小轴 + 小 K + 排序"复合场景。

完整 key 映射表:

数据类型 普通 Bitonic MergeSort Bitonic+MergeSort
INT8 1001 51001 - -
INT16 1002 51002 - -
INT32 1003 51003 - -
INT64 1004 51004 - -
UINT8 2001 52001 - -
UINT16 2002 52002 - -
UINT32 2003 52003 - -
UINT64 2004 52004 - -
FLOAT 3003 53003 13003 63003
FLOAT16 3002 53002 13002 63002
BF16 4002 54002 14002 64002

整数类型不支持 Merge Sort(optDataTypeBitMap 仅包含 FLOAT/FLOAT16/BF16)。FP32 MoreCore/IntraCore 是独立算法路径,不叠加 Bitonic 偏移。

3.3 IS_BITONIC_SORT 模板参数

Kernel 侧通过编译时模板参数 IS_BITONIC_SORTbool)控制双调路径的启用:

// 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>(...);
#endif

IS_BITONIC_SORT 传递到 7 类算子模板,覆盖 TopKV2 全部计算路径:

模板 文件 IS_BITONIC_SORT 作用
RadixSortTopK radix_sort_top_k.h SortTopKRes 后调用 RunBitonicSmallTopKFinalize
RadixSortTopKSingleBlock radix_sort_top_k_single_block.h ProcessSingleTime 后调用 RunBitonicSmallTopKFinalize
RadixSortTopKSingleCore radix_sort_top_k_single_core.h SortTopKRes 后调用 RunBitonicSmallTopKFinalizeIsUseRadixGather 编译时短路
RadixSortTopKMultiCoreOptimization radix_sort_top_k_inter_core_template_optimization.h ProcessSingleLoop 后调用 RunBitonicSmallTopKFinalize
SortAndTopKMoreCore sort_and_top_k_more_core.h GetBitonicSmallTopKRes 独立路径替代 GetTopKRes
MergeSort top_k_merge_sort.h VbsMergeSort 后调用 RunBitonicSmallMergeSortFinalizeNonLast
TopKNonLastSmallAxisNonTranspose top_k_non_last_small_non_transpose.h RunTopK/RunMergeSort 后调用 RunBitonicSmallTopKFinalizeNonLast

3.4 编译时优化

原实现中 IsUseRadixGather 通过运行时成员变量 sortPolicy_ 判断是否启用 Gather 路径。重构后:

  • IS_BITONIC_SORT = true 时,IsUseRadixGatherif constexpr 中编译期短路返回
  • Host 侧 IsBitonicSmallTopkMode 已保证 IS_BITONIC_SORT = true 蕴含 IS_SORTk ∈ [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 文件整体架构

top_k_small_size_bitonic_sort.h
├── 常量定义
│   ├── BITONIC_SMALL_TOPK_SIZE = 32        (双调网络规模)
│   ├── BITONIC_SMALL_TOPK_MAX_ROWS = 32    (SIMT 多行并行上限)
│   └── BITONIC_SMALL_TOPK_THREADS = 1024   (SIMT 最大线程数)
│
├── 类型萃取
│   ├── BitonicSmallGatherIndexType<T>      (Gather 索引类型选择)
│   ├── BitonicSmallGatherSignedIndexType<T>(有符号 Gather 索引)
│   ├── BitonicSmallRegType<T>              (SIMT 标量比较类型提升)
│   └── IsBitonicFloatType<T>               (浮点类型判定)
│
├── Reg 路径 (SIMD 向量寄存器)
│   ├── BitonicSmallRegLoadB16Bits          (16位→32位位模式展开)
│   ├── BitonicSmallRegStoreB16Bits         (32位→16位位模式打包)
│   ├── BitonicSmallRegBuildKey32           (值→可比较无符号key)
│   ├── BitonicSmallReg32SwapStage          (32位compare-swap)
│   ├── BitonicSmallReg32BitonicNetwork     (32位完整双调网络)
│   ├── BitonicSmallRegFinalizeSelectionB16 (16位收尾)
│   ├── BitonicSmallReg64SwapStage          (64位compare-swap)
│   ├── BitonicSmallReg64BitonicNetwork     (64位完整双调网络)
│   ├── BitonicSmallRegFinalizeSelectionB64 (64位收尾)
│   ├── BitonicSmallRegValueBetter          (值比较含NaN语义)
│   ├── BitonicSmallRegValueEquivalent      (值等价判定)
│   ├── BitonicSmallRegSwapStage            (通用compare-swap)
│   ├── BitonicSmallRegBitonicNetwork       (通用完整双调网络)
│   └── BitonicSmallRegFinalizeSelection    (Reg路径总分发)
│
├── SIMT 路径 (Warp 级标量线程)
│   ├── BitonicSmallValueBetter             (标量值比较)
│   ├── BitonicSmallValueEquivalent         (标量值等价)
│   ├── BitonicSmallInvalidIndex            (无效索引标记)
│   ├── BitonicSmallSwapStage               (warp级compare-swap)
│   ├── BitonicSmallBitonicNetwork          (warp级完整网络)
│   ├── BitonicSmallAllEqualSwapStage       (全等场景compare-swap)
│   ├── BitonicSmallAllEqualBitonicNetwork  (全等场景网络)
│   ├── BitonicSmallNthSetBit              (位掩码第rank个置位)
│   ├── BitonicSmallFinalizeSelection       (已排序候选收尾)
│   └── BitonicSmallFinalizeExactSelection  (未排序候选收尾)
│
├── SIMT Kernel (asc_vf_call 调度)
│   ├── BitonicGatherThresholdTileKernel    (阈值gather收集)
│   ├── BitonicFinalizeExactSelectionKernel (精确选择收尾)
│   ├── BitonicFinalizeSmallSourceRowsKernel(小源行收尾)
│   └── BitonicFinalizeSelectionRowsKernel  (行批量收尾)
│
└── 对外入口
    ├── RunBitonicGatherThresholdTile       (gather封装)
    ├── RunBitonicFinalizeExactSelection    (精确选择封装)
    ├── RunBitonicFinalizeSmallSourceRows   (小源行封装)
    ├── RunBitonicFinalizeSelectionRows     (行批量封装)
    ├── RunBitonicSmallTopKFinalize         (尾轴主入口)
    ├── RunBitonicSmallTopKFinalizeNonLast  (非尾轴TopK入口)
    └── RunBitonicSmallMergeSortFinalizeNonLast (非尾轴MergeSort入口)

4.2 两条实现路线

路径 指令层 标注 适用场景 关键 API
Reg 路径 SIMD 向量寄存器 __simd_callee__ 非 1 字节类型 + indices 为 uint32_t Reg::Gather, Reg::Select, Reg::Compare
SIMT 路径 Warp 级标量线程 __simt_callee__ / __simt_vf__ 1 字节类型或 indices 为 int64_t asc_shfl_xor, asc_ballot, __popc

1 字节类型不能走 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)

双调网络用整数比较替代浮点/有符号数比较,需将各种类型的数据位模式映射为可单调比较的 uint32 key:

数据类型 映射规则 说明
无符号整数 key = rawBits 天然单调
有符号整数 key = rawBits XOR signBit 翻转符号位,使负数排在正数前
浮点(正数) key = rawBits XOR signBit 仅翻转符号位
浮点(负数) key = rawBits XOR allBits 全位取反,使绝对值大的负数排更后
TopK 最小值 key XOR allBits 整体取反,反转排序方向

NaN 特殊处理:NaN 的指数全 1 且尾数非 0。IsLargest 时 NaN key 设为 allBits(排在最后),否则设为 1(排在最前),保证 NaN 总在结果末尾。浮点 -0 会被清零,与 +0 统一处理。

4.4 双调排序网络

32-lane 双调网络由 5 级归并构成,每级内部按递减 stride 做 compare-swap,共 15 个 SwapStage:

Stage 1 (size=2):    (stride=1)
Stage 2 (size=4):    (stride=2), (stride=1)
Stage 3 (size=8):    (stride=4), (stride=2), (stride=1)
Stage 4 (size=16):   (stride=8), (stride=4), (stride=2), (stride=1)
Stage 5 (size=32):   (stride=16), (stride=8), (stride=4), (stride=2), (stride=1)

先构建双调序列(size 递增),再做双调合并(stride 递减),最终全排序。

4.4.1 Reg 路径 Compare-Swap(BitonicSmallReg32SwapStage)

每个 lane 执行以下步骤:

  1. 计算对端 lane:peerLane = lane XOR stride
  2. Reg::Gather 获取对端 value/key/index/group
  3. 按本地/对端分为 low/high 两部分
  4. 比较判定:CompareGroup=true 时按 (group, index) 字典序;false 时按 key 大小
  5. 计算交换掩码:swapMask = (compare && lowValid) || !highValid
    • 低位比高位"差"且低位有效 → 交换
    • 高位无效 → 强制交换(把无效项挤到低位)
  6. 计算方向掩码:Size==32 时全升序;否则按子网方向
  7. 决定取值:takePeerMask = !(swapMask XOR directionMask)

4.4.2 SIMT 路径 Compare-Swap(BitonicSmallSwapStage)

每个 lane 线程持有一个标量元素:

  1. asc_shfl_xor 获取对端 lane 的 value/index
  2. 标量比较 + 条件赋值实现交换
  3. 交换逻辑与 Reg 路径一致,但用标量条件判断替代掩码运算

4.5 候选选择与收尾排序

4.5.1 已排序候选收尾(BitonicSmallFinalizeSelection)

处理 TopK/RadixSelect 输出的已按值排序的候选集:

输入: k 个已排序候选 (lane 0~k-1 有效, lane k~31 无效)
  │
  ├─ 取阈值: threshold = 第 k-1 个元素的值 (已排序,直接 shfl 取)
  │
  ├─ 全等快速路径: 若首元素 == 阈值 → 所有候选相等
  │   └─ 仅按 index 排序 (AllEqualBitonicNetwork), 直接返回
  │
  ├─ 重复值检测: 检查是否存在与前驱相等的元素
  │   └─ 无重复 → 已是有序序列, 直接返回
  │
  └─ BITONIC 兼容性重构:
      ├─ 步骤1: 按 index 排序恢复原始 lane 序 (BitonicNetwork<false>)
      ├─ 步骤2: 计算 strict/equal 掩码
      │   strict = 值优于阈值, equal = 值等于阈值
      ├─ 步骤3: 用 NthSetBit 找到每个输出 lane 的源 lane
      │   sourceLane = NthSetBit(strictMask/equalMask, rank)
      ├─ 步骤4: asc_shfl 按源 lane 重新收集数据
      └─ 步骤5: 按值排序 (BitonicNetwork<true>), 输出最终结果

4.5.2 未排序候选收尾(BitonicSmallFinalizeExactSelection)

处理 Gather 路径输出的未排序精确候选集。与上述区别在于候选集未按值排序,需要先归约求阈值:

输入: k 个未排序候选
  │
  ├─ 蝶形归约求阈值:
  │   for stride = 16, 8, 4, 2, 1:
  │       peer = asc_shfl_xor(threshold, stride)
  │       if ValueBetter(threshold, peer): threshold = peer
  │   → 得到候选中的最差值作为阈值
  │
  └─ 其余 BITONIC 兼容性重构逻辑与 FinalizeSelection 相同

4.5.3 Reg 路径收尾(BitonicSmallRegFinalizeSelection)

Reg 路径的收尾逻辑与 SIMT 路径语义等价,但用向量寄存器指令实现:

  1. 加载 value 和 index,越界 index 填 UINT32_MAX 标记无效
  2. 构建排序 key(BitonicSmallRegBuildKey32
  3. 取第 k 个元素的 key 作为阈值,分三组:
    • strict (key > 阈值) → group=0
    • 其余 valid → group=1
    • 无效 → group=2
  4. 先按 (group, index) 排序(BitonicNetwork<true>),再按 key 全排序(BitonicNetwork<false>
  5. 存回 UB

sizeof(T) 分发到不同特化路径:

sizeof(T) 路径 说明
8 字节 (int64/uint64) BitonicSmallRegFinalizeSelectionB64 拆双 32 位寄存器
2 字节 (half/bf16/int16/uint16) BitonicSmallRegFinalizeSelectionB16 位模式展开为 uint32
4 字节 (int32/uint32/float) 通用路径 直接按 value 比较
1 字节 (int8/uint8) 不走 Reg,回退 SIMT 硬件限制

4.6 阈值 Gather 收集(BitonicGatherThresholdTileKernel)

针对大轴场景,使用 SIMT warp 内 ballot + popcount 做前缀和,按阈值收集候选元素:

两趟扫描:
  pass 0: 收集 key < threshold (严格优于阈值) 的元素
  pass 1: 收集 key == threshold (等于阈值) 的元素

每趟内:
  for base = 0; base < axisSize; base += 32:
    col = base + lane
    take = inRange && (key[col] < threshold 或 == threshold)
    takeMask = asc_ballot(take)           // warp 内收集哪些 lane 选中
    rank = __popc(takeMask & lowerMask)   // 前缀和: 本 lane 之前有多少选中
    outputPos = written + rank
    if (take && outputPos < quota):
        output[outputPos] = input[col]    // 写入对应位置
    written += __popc(takeMask)           // 更新已写入计数

4.7 多行批量调度

SIMT 路径支持多行并行处理,threadIdx.y 为行索引,threadIdx.x 为 lane (0~31):

函数 用途 每批最大行数
RunBitonicFinalizeSelectionRows 行批量收尾排序 BITONIC_SMALL_TOPK_MAX_ROWS (32)
RunBitonicFinalizeSmallSourceRows 小源行收尾(axisLen ≤ 32) BITONIC_SMALL_TOPK_MAX_ROWS (32)

分批调度:for (rowStart = 0; rowStart < rowCount; rowStart += 32)

5. 模板集成点详解

5.1 尾轴场景

RadixSortTopK / RadixSortTopKSingleCore / RadixSortTopKSingleBlock

这三类模板在 AscendC::SortAscendC::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::TopKAscendC::Sort 高阶 API 后,调用专用收尾入口:

// top_k_non_last_small_non_transpose.h:259
RunBitonicSmallTopKFinalizeNonLast<SortT, IsLargest, IS_BITONIC_SORT>(
    topkValue_, topkIndex_, this->sortInput_, this->axisLen_, kValue_, ...);

该入口的分发逻辑:

sizeof(SortT) axisLen 路径
1 ≤ 32 RunBitonicFinalizeSmallSourceRows(小源行)
1 > 32 RunBitonicFinalizeSelectionRows(SIMT 行批量)
≠ 1 任意 BitonicSmallRegFinalizeSelection(Reg 逐行)

MergeSort

// top_k_merge_sort.h:274
RunBitonicSmallMergeSortFinalizeNonLast<SortT, IsLargest, IS_BITONIC_SORT>(dst, dstIndex, ...);

分发逻辑:

sizeof(SortT) 路径
> 4 或 = 1 RunBitonicFinalizeSelectionRows(SIMT)
2 或 4 BitonicSmallRegFinalizeSelection(Reg)

64 位类型在 merge sort 场景走 SIMT 路径,避免 Reg 路径的双寄存器拆分复杂度。

6. 数据类型支持矩阵

数据类型 sizeof Reg 路径 SIMT 路径 说明
INT8 1 不支持 支持 Reg API 最小 16 位
UINT8 1 不支持 支持 同上
INT16 2 支持 (B16) 支持 位模式展开为 uint32
UINT16 2 支持 (B16) 支持 同上
HALF 2 支持 (B16) 支持 同上
BF16 2 支持 (B16) 支持 同上
INT32 4 支持 (通用) 支持 直接按值比较
UINT32 4 支持 (通用) 支持 同上
FLOAT 4 支持 (通用) 支持 同上
INT64 8 支持 (B64) 支持 拆双 32 位寄存器
UINT64 8 支持 (B64) 支持 同上

7. 流水线同步设计

7.1 SortAndTopKMoreCore 逐行处理同步

GetBitonicSmallTopKRes 逐行处理时,valuesLocal 来自 depth=1 队列,indicesLocal 来自固定 TBuf。行间需防止 MTE3 读 UB 与下一行 MTE2 写 UB 的竞争:

Row N:   [MTE2写入UB] →MTE2_V→ [V计算] →V_MTE3→ [MTE3读UB→GM] →MTE3_MTE2→ FreeTensor
                                                                    ↓ 等待MTE3完成
Row N+1: [MTE2写入同一UB] ← AllocTensor(MTE3已安全完成)
// 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 两轮排序策略

收尾排序采用两轮双调网络:

  1. 第一轮CompareGroup=true):按 (group, index) 排序,恢复原始 lane 顺序
  2. 第二轮CompareGroup=false):按值排序,输出最终有序结果

这是为了兼容 TopK 高阶 API 的输出语义——当存在重复值时,TopK 选择的候选可能与 Bitonic gather 的候选不同,需要先恢复原始顺序再重新排序。

8.4 IS_BITONIC_SORT 编译时参数

选择编译时模板参数而非运行时变量,使得:

  • if constexpr 在编译期消除非 Bitonic 路径的代码
  • IsUseRadixGather 等条件判断编译期短路,零运行时开销
  • 减少了 tiling data 中 sortPolicy 字段的读取和存储

9. 文件依赖关系

top_k_v2_apt.cpp (入口分发)
  ├── radix_sort_top_k.h
  │   └── top_k_small_size_bitonic_sort.h ★
  ├── radix_sort_top_k_single_block.h
  │   └── top_k_small_size_bitonic_sort.h ★
  ├── radix_sort_top_k_single_core.h
  │   └── top_k_small_size_bitonic_sort.h ★
  ├── radix_sort_top_k_inter_core_template_optimization.h
  │   └── top_k_small_size_bitonic_sort.h ★
  ├── sort_and_top_k_more_core.h
  │   └── top_k_small_size_bitonic_sort.h ★
  ├── top_k_merge_sort.h
  │   └── top_k_small_size_bitonic_sort.h ★
  └── top_k_non_last_small_non_transpose.h
      └── top_k_small_size_bitonic_sort.h ★

top_k_small_size_bitonic_sort.h (★核心实现)
  ├── top_k_constant_var_simd.h  (常量定义)
  └── top_k_util_type_simd.h     (IsUseRadixGather)
likedislike
cy_hwcy_hw成员
12 天前 添加了label:requirement
cy_hwcy_hw成员
12 天前 关联了pull request:TopkV2算子支持小k场景下的双调排序
陈思
陈思成员
11 天前 评论:
likedislike
CANN-robotCANN-robot成员
11 天前 将 caoyan_huawei 设为负责人
cy_hwcy_hw成员
10 天前 修改了issue 的描述
cy_hwcy_hw成员
9 天前 修改了issue 的描述
cy_hwcy_hw成员
9 天前 修改了issue 的描述
CANN-robotCANN-robot成员
7 天前 关闭了 issue
CANN-robotCANN-robot成员
7 天前 添加了label:resolved