Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
argmax算子在 950DT中,出现了同带宽性能不如910B和H800的情况,主要围绕bank冲突,小shape逻辑,NDDMA带宽无提升(相比950PR)几个问题展开,具体方案如下
性能提升
此 PR 主要对 arg_max_with_value 算子的 tiling 策略和 kernel 实现进行了多方面优化:引入动态核数机制(按每 4KB 数据分配一核)、新增 originCoreNum_ 以保持 GROUP_REDUCE 全核同步的自洽、移除过时的 ARA MODE5/105 和 MODE6/106 分支并统一到 MODE4、为 ARA_CUT_A_AND_NEXT_A 新增 flat-copy 快速路径、替换 ArgMaxRaInt64/ArgMaxAraInt64 为支持非对齐源/目标步长的 ArgMaxAraUnAlignInt64/ArgMaxAraUnAlign 系列新 kernel 原语,以及增加 AR_GATHER 的 bank 冲突规避逻辑。此外包含大量代码风格统一(缩进、指针格式、花括号位置)。
arg_max_with_value
originCoreNum_
ARA_CUT_A_AND_NEXT_A
ArgMaxRaInt64
ArgMaxAraInt64
ArgMaxAraUnAlignInt64
ArgMaxAraUnAlign
AR_GATHER
在 ArgCommonBaseTiling 中新增 originCoreNum_ 保存物理核数,SetShapeInfo() 按每 4KB 输入数据动态计算 coreNum_(MAX_SIZE_USING_SINGLE_CORE);GROUP_REDUCE 的核数判断、分核及 workspace 大小计算全链路使用 originCoreNum_,确保 SyncAll 同步自洽。
ArgCommonBaseTiling
SetShapeInfo()
coreNum_
MAX_SIZE_USING_SINGLE_CORE
GROUP_REDUCE
SyncAll
在 CalcSplitInfoForArA() 中删除 MODE5(切 nextA)和 MODE6(保底 R 轴动态切分)分支,统一由 MODE4 的 CalcCutRA() 处理;kernel 侧 arg_max_with_value_ara.h 同步移除对应的 ProcessWithCutR()、ProcessWithCutNextA() 及相关辅助函数和常量。
CalcSplitInfoForArA()
CalcCutRA()
arg_max_with_value_ara.h
ProcessWithCutR()
ProcessWithCutNextA()
在 CalcSplitInfoForArACutAAndNextA() 中优先尝试 flat 路径——当 nextA 对齐后整个 UB 需求在可用 UB 范围内时,调用新增的 CalcSplitInfoForArACutAAndNextAFlat() 设置 useFlatPath_=1;kernel 侧 arg_max_with_value_ara_cut_a_and_next_a.h 新增 ProcessFlat() 和 CopyInXFlat(),整块搬入后一次 ComputeVF 完成。
CalcSplitInfoForArACutAAndNextA()
CalcSplitInfoForArACutAAndNextAFlat()
useFlatPath_=1
arg_max_with_value_ara_cut_a_and_next_a.h
ProcessFlat()
CopyInXFlat()
ComputeVF
在 arg_max_with_value_base.h 中新增 ArgMaxAraUnAlign、ArgMaxAraUnAlignCompact 和 ArgMaxAraUnAlignInt64 三个模板函数,支持独立的源/目标 dimA1 步长(dimA1Src/dimA1Dst);其中 ArgMaxAraUnAlignCompact 在 dimR ≥ 2 且 dimA1Src × 2 ≤ VL 时通过两行打包合并和寄存器 Gather 实现紧凑比较。arg_max_with_value_ra.h、arg_max_with_value_ara_gather.h、arg_max_with_value_group_reduce.h 中的 int64 路径统一切换至新原语。
arg_max_with_value_base.h
ArgMaxAraUnAlignCompact
dimA1
dimA1Src
dimA1Dst
dimR ≥ 2
dimA1Src × 2 ≤ VL
Gather
arg_max_with_value_ra.h
arg_max_with_value_ara_gather.h
arg_max_with_value_group_reduce.h
在 SetShapeInfoHighPerf() 中为 AR_GATHER 增加条件 (R × eleBytes) % BLOCK_SIZE != 0,排除步长 32B 对齐导致 UB bank 冲突的场景;同时放宽 ARA_CUT_A_AND_NEXT_A 中 nextA 不切核时的最小字节门槛至 64B(int64 除外),并调整 CalcSplitInfoForArA() 中 MODE1/MODE2 的 MAX_NEXTA_SIZE 阈值为 2 倍。
SetShapeInfoHighPerf()
(R × eleBytes) % BLOCK_SIZE != 0
MAX_NEXTA_SIZE
/assign @RoyHys
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
一、背景信息 (必填)
argmax算子在 950DT中,出现了同带宽性能不如910B和H800的情况,主要围绕bank冲突,小shape逻辑,NDDMA带宽无提升(相比950PR)几个问题展开,具体方案如下
二、价值/作用 (必填)
性能提升
三、设计方案 (必填)
变更摘要
此 PR 主要对
arg_max_with_value算子的 tiling 策略和 kernel 实现进行了多方面优化:引入动态核数机制(按每 4KB 数据分配一核)、新增originCoreNum_以保持 GROUP_REDUCE 全核同步的自洽、移除过时的 ARA MODE5/105 和 MODE6/106 分支并统一到 MODE4、为ARA_CUT_A_AND_NEXT_A新增 flat-copy 快速路径、替换ArgMaxRaInt64/ArgMaxAraInt64为支持非对齐源/目标步长的ArgMaxAraUnAlignInt64/ArgMaxAraUnAlign系列新 kernel 原语,以及增加AR_GATHER的 bank 冲突规避逻辑。此外包含大量代码风格统一(缩进、指针格式、花括号位置)。主要改动
1. 动态核数与
originCoreNum_分离在
ArgCommonBaseTiling中新增originCoreNum_保存物理核数,SetShapeInfo()按每 4KB 输入数据动态计算coreNum_(MAX_SIZE_USING_SINGLE_CORE);GROUP_REDUCE的核数判断、分核及 workspace 大小计算全链路使用originCoreNum_,确保SyncAll同步自洽。2. 移除过时的 ARA MODE5/105 和 MODE6/106
在
CalcSplitInfoForArA()中删除 MODE5(切 nextA)和 MODE6(保底 R 轴动态切分)分支,统一由 MODE4 的CalcCutRA()处理;kernel 侧arg_max_with_value_ara.h同步移除对应的ProcessWithCutR()、ProcessWithCutNextA()及相关辅助函数和常量。3.
ARA_CUT_A_AND_NEXT_A新增 flat-copy 快速路径在
CalcSplitInfoForArACutAAndNextA()中优先尝试 flat 路径——当 nextA 对齐后整个 UB 需求在可用 UB 范围内时,调用新增的CalcSplitInfoForArACutAAndNextAFlat()设置useFlatPath_=1;kernel 侧arg_max_with_value_ara_cut_a_and_next_a.h新增ProcessFlat()和CopyInXFlat(),整块搬入后一次ComputeVF完成。4. 新 kernel 原语
ArgMaxAraUnAlign系列在
arg_max_with_value_base.h中新增ArgMaxAraUnAlign、ArgMaxAraUnAlignCompact和ArgMaxAraUnAlignInt64三个模板函数,支持独立的源/目标dimA1步长(dimA1Src/dimA1Dst);其中ArgMaxAraUnAlignCompact在dimR ≥ 2且dimA1Src × 2 ≤ VL时通过两行打包合并和寄存器Gather实现紧凑比较。arg_max_with_value_ra.h、arg_max_with_value_ara_gather.h、arg_max_with_value_group_reduce.h中的 int64 路径统一切换至新原语。5.
AR_GATHERbank 冲突规避与阈值调整在
SetShapeInfoHighPerf()中为AR_GATHER增加条件(R × eleBytes) % BLOCK_SIZE != 0,排除步长 32B 对齐导致 UB bank 冲突的场景;同时放宽ARA_CUT_A_AND_NEXT_A中 nextA 不切核时的最小字节门槛至 64B(int64 除外),并调整CalcSplitInfoForArA()中 MODE1/MODE2 的MAX_NEXTA_SIZE阈值为 2 倍。