已关闭
[Requirement|需求建议]: sinkhorn support 950 #2788
Davon创建于  8月21日关闭于  29 天前
Davon成员
8月21日 创建

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

一、背景信息 (必填)

Sinkhorn 算子是 CANN ops-math 仓中的 Sinkhorn 最优传输距离计算算子,对二维成本矩阵 cost(R, C) 迭代执行 Sinkhorn-Knopp 矩阵缩放,得到最优传输计划 p(R, C)。典型应用为 MoE 大模型专家路由、最优传输/分布匹配、点云配准等场景。
当前算子仅支持 Atlas A2(ascend910b)和 Atlas A3(ascend910_93),不支持 Ascend 950PR/950DT。本次需求将 Sinkhorn 算子适配到 950 平台,并在适配过程中修复 FP16 半精度计算精度不足、BF16 搬运字节数错误、跨 tile 缺少同步、NaN 死循环等已知问题。

二、价值/作用 (必填)

使 Sinkhorn 算子在 Ascend 950PR/950DT 上可用,支撑 950 平台上的 MoE 专家路由和最优传输场景。
修复 FP16 输入精度问题:将 FP16 输入从全程半精度计算改为内部 FP32 计算,消除迭代精度累积误差。
修复 BF16 搬运错误和跨 tile 同步缺失,确保大矩阵多核场景下结果正确。
修复 NaN 导致的 kernel 死循环,提升极端输入下的鲁棒性。

三、设计方案 (必填)

3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)

Aclnn 直调 aclnnSinkhorn / aclnnSinkhornGetWorkspaceSize 两段式接口

3.2 总体设计
3.2.1 算子支持的数据类型

cost FLOAT / FLOAT16 / BFLOAT16
p 同 cost(Follow)
tol FLOAT(aclScalar)

3.2.2 host侧设计

OpDef:新增 this->AICore().AddConfig("ascend950"),使算子在 950 上可编译。
Tiling:FP16 的 sizeOfDataType 从 sizeof(uint16_t) 改为 sizeof(float)(因为内部按 FP32 计算,workspace 按 FP32 大小分配)。新增 pFloatGlobal workspace:totalRow * totalCol * sizeof(float),用于 FP32 存储 exp(cost)。
aclnn 校验增强:新增 CheckTolValid(tol 必须为正有限 float)、MAX_COST_ROW=10000(行数上限)、行列正值校验。

3.2.3 kernel侧设计

FP16 升精度:

  • dispatch 从 KernelSinkhorn<half, half> 改为 KernelSinkhorn<float, half>
  • 新增 pFloatGlobal(FP32 workspace):exp(cost) 始终以 FP32 保存,不再降回输入 dtype
  • CopyOutForExp 从写 pGlobal 改为写 pFloatGlobal(FP32)
  • CopyInFromP 从读 pGlobal 改为读 pFloatGlobal(FP32)
  • 删除 FP16/BF16 的 CopyInFromP 和 CopyOutForExp 特化(不再需要输入侧 Cast)
    950 VF Cast 指令(#if NPU_ARCH == 3510):
  • 新增 SinkhornHalfToFloatVF / SinkhornFloatToHalfVF,使用 Reg::Cast SIMD 指令
  • FloatToHalf 用 CAST_RINT(四舍五入)替代原来的 CAST_TRUNC(截断),精度更优
  • A2/A3 走 #else 分支,使用标准 Cast() API
    BF16 搬运修复:
  • SaveP<bfloat16_t> 中 blockLen 从 sizeof(float) 改为 sizeof(bfloat16_t)(按实际宽度搬运)
  • 一次性 Cast + DataCopyPad 改为逐行处理,避免 950 上 stride/对齐不匹配
3.3 支持硬件

Ascend 950PR / Ascend 950DT
Atlas A3 训练/推理
Atlas A2 训练/推理

3.4 算子约束限制

cost 必须为二维矩阵(rank=2),行数 ≤ 10000,列数 ≤ 4096。
tol 必须为正有限 FLOAT 值(空指针默认 0.0001)。
cost 和 p 的 shape 必须完全一致。
cost 支持非连续 tensor,p 不支持非连续 tensor。
cost 建议值域 0, 1(exp(cost) 对大值可能溢出为 inf,kernel 会通过 NaN 安全收敛判断退出循环,但输出可能包含 NaN)。
迭代次数由 tol 控制,无最大迭代次数上限(非 NaN 场景 Sinkhorn 数学上保证收敛)。
BF16 支持依赖 SoC 版本:950 和 910B~910E 支持 BF16,其他平台仅 FLOAT/FLOAT16。
950 上 FP16/BF16 的 Cast 使用 VF 寄存器指令(NPU_ARCH == 3510),A2/A3 使用标准 Cast() API。
💡 备注(选填)

likedislike
DDavon成员
8月21日 添加了label:requirement
陈思
陈思成员
8月21日 评论:

/assign @Davon14272

likedislike
CANN-robotCANN-robot成员
8月21日 将 Davon14272 设为负责人
CANN-robotCANN-robot成员
29 天前 关闭了 issue
CANN-robotCANN-robot成员
29 天前 添加了label:resolved