已合并
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
已合并
HKFLYE创建于 8月17日
HKFLYE
HKFLYE成员
8月17日

描述

本 PR 适配 Addmm、Baddbmm 及 AddmmWeightNz 的 ACLNN 接口路由,使 A2/A3(DAV_2201)上的 broadcast bias 场景能够直接使用 GemmV3 完成矩阵乘、alpha/beta 缩放及 bias broadcast。

主要修改如下:

1. 清理 broadcast bias 路由限制

  • 删除 CheckAddmmTensorShapeNeedBroadcast。
  • 简化 CheckCubeMathTypeForAddMm,仅保留 cubeMathType 取值及平台能力校验。
  • 不再因为 bias shape 与矩阵乘输出 shape 不一致而拒绝 cubeMathType=USE_FP32_ADD。
  • 路由限制在 A2/A3(DAV_2201),不改变其他架构的执行逻辑。

2. 调整 Addmm 和 Baddbmm 的 GemmV3 路由

  • cubeMathType=USE_FP32_ADD 时,允许 broadcast bias 进入 GemmV3。
  • 区分“接口计算需要 FP32 中间结果”和“GemmV3 实际输出 FP32”,避免 FP16/BF16 输出错误使用 16in32out GemmV3。

3. 调整 AddmmWeightNz 的 GemmV3 路由

  • 取消 A2/A3 上将 cubeMathType=USE_FP32_ADD 转换为 cubeMathType=KEEP_DTYPE 的逻辑。
  • cubeMathType=USE_FP32_ADD 时,支持 AddmmWeightNz 直接路由至 GemmV3。
  • 支持 broadcast bias 与 ND×NZ 矩阵输入组合进入 GemmV3。
  • 支持 FP16/BF16 输入及对应的低精度输出。
  • 保留原有 16in32out 场景,支持 GemmV3 输出 FP32。

4. 补齐 GemmV3 OpDef 类型声明

补充以下四组类型和格式组合:

A 输入 B 输入 Bias 输入 输出
FP16 ND FP16 ND FP16 ND FP32 ND
BF16 ND BF16 ND BF16 ND FP32 ND
FP16 ND FP16 NZ FP16 ND FP16 ND
BF16 ND BF16 NZ BF16 ND BF16 ND

前两组用于 ND×ND 低精度 bias 的 16in32out 场景,后两组用于 AddmmWeightNz 在 cubeMathType=USE_FP32_ADD 下保持低精度输出的场景。

5. 补齐 A2/A3 GemmV3 binary 配置

在以下配置中补充 ND×NZ 同类型输出 Kernel:

  • ascend910b/gemm_v3_binary.json
  • ascend910_93/gemm_v3_binary.json

新增配置:

Binary 输入输出组合
GemmV3_ND_NZ_ND_ND_FP16_FP16_FP16_FP16 FP16 ND×NZ、FP16 bias、FP16 输出
GemmV3_ND_NZ_ND_ND_BF16_BF16_BF16_BF16 BF16 ND×NZ、BF16 bias、BF16 输出

未修改 ascend950 和 ascend350 的 binary 配置。

6. 支持的 bias broadcast shape

Addmm/AddmmWeightNz

输出 shape 为 [M,N]:

Broadcast 类型 Bias shape
N broadcast [N]、[1,N]
M broadcast [M,1]
Scalar broadcast [1]、[1,1]
非 broadcast [M,N]

Baddbmm

输出 shape 为 [B,M,N]:

Broadcast 类型 Bias shape
N broadcast [N]、[1,N]、[1,1,N]
M broadcast [M,1]、[1,M,1]
Scalar broadcast [1]、[1,1]、[1,1,1]
Per-batch scalar [B,1,1]
非 broadcast [B,M,N]

关联的Issue

Issue #4844

测试

1. ATK 精度泛化用例

在 A2/A3 环境执行2000条 ATK L0 双标杆精度泛化用例:

接口 用例数 Bias shape 分布
Addmm 400 N:134,M:67,scalar:133,full:66
InplaceAddmm 400 full:400
AddmmWeightNz 400 N:134,M:67,scalar:133,full:66
Baddbmm 400 N:135,M:89,scalar:176
InplaceBaddbmm 400 full:400

精度用例覆盖:

  • FP16/BF16 矩阵输入;
  • FP16、BF16及FP32输出;
  • ND×ND 和 ND×NZ;
  • cubeMathType=USE_FP32_ADD;
  • FP16/BF16 16in32out;
  • transA/transB 四种组合;
  • N、M、scalar 和 per-batch scalar broadcast;
  • 对齐与非对齐 shape;
  • 大小 shape 及尾块场景。

双标杆精度首轮结果为 1998/2000 通过,两条失败用例定点复检均通过,未发现稳定复现的精度问题。

2. 性能验证

执行100条三通路性能对比用例:

分类 用例数
Addmm 50
Baddbmm 50
N broadcast 33
M broadcast 33
Scalar broadcast 34

对比以下三条通路:

  1. GemmV3 + broadcast bias;
  2. GemmV3 + 等价 full-shape bias;
  3. ATB PPMatMul + Muls + Axpy。

GemmV3 与 PPMatMul 使用相同的原始 shape、物理布局、transpose 和 contiguous 条件。

性能结果:

  • 小算子链/GemmV3 耗时比中位数:2.2852;
  • 比值范围:1.1051~2.8523;
  • GemmV3 broadcast 不慢于小算子拼接:100/100;
  • 三条通路输出一致性检查:100/100通过。

文档更新

更新以下 ACLNN 接口文档:

  • aclnnAddmm&aclnnInplaceAddmm.md
  • aclnnBaddbmm&aclnnInplaceBaddbmm.md
  • aclnnAddmmWeightNz.md

文档更新内容包括:

  • A2/A3 上 cubeMathType=USE_FP32_ADD 对 broadcast bias 的支持说明;
  • AddmmWeightNz 对 cubeMathType=USE_FP32_ADD 的支持说明;
  • AddmmWeightNz 支持的 [M,1]、[1,1] 等 broadcast shape;
  • FP16/BF16 输入、低精度输出和 FP32 中间计算行为说明。
  • GemmV3 broadcast bias 的 Kernel 设计和精度、性能验证方法由对应设计文档及测试 README 说明。

类型标签

AI/Agent生成声明

likedislike
Pull Request已成功合入, 合并人@CANN-robot
(感谢 HKFLYE 的贡献)
HKFLYEHKFLYE成员
8月17日 创建了 pull request,commit f46d5dc6
atomgit-bot
atomgit-bot
8月17日 评论:

变更摘要

本 PR 主要调整 A2/A3(DAV_2201)上 aclnnAddmm、aclnnBaddbmm、aclnnAddmmWeightNz 的 ACLNN 路由逻辑,使 cubeMathType=USE_FP32_ADD 且存在 broadcast bias 的场景能够直接路由至 GemmV3 完成矩阵乘、alpha/beta 缩放和 bias 广播,同时补齐相应 OpDef 类型声明与 A2/A3 的 binary Kernel 配置,并清理了原有的 broadcast 限制校验。

主要改动

  • 清理 broadcast bias 路由限制:在 cube_util.cpp/cube_util.h 中删除 CheckAddmmTensorShapeNeedBroadcast,并将 CheckCubeMathTypeForAddMm 简化为仅校验 cubeMathType 取值和平台能力,不再因 bias 与矩阵乘输出 shape 不一致而拒绝 USE_FP32_ADD。
  • 调整 Addmm/Baddbmm 的 GemmV3 路由:aclnn_addmm.cpp 与 aclnn_baddbmm.cpp 中移除 needBroadcast 判断,改为校验 bias dtype 与矩阵输入 dtype 一致或为 DT_FLOAT,使 USE_FP32_ADD 和 16in32out 场景可直接走 ExecGemmV3WithAlphaBetaOp;同时区分 enableFp32Output 与 enableGemm16In32Out,避免 FP16/BF16 输出误用 16in32out。
  • 放开 AddmmWeightNz 的 USE_FP32_ADD 路由:aclnn_addmm.cpp 中删除 routeCubeMathType4ToCubeMathType0DAV_2201 调用,允许 cubeMathType=USE_FP32_ADD 时通过 CheckGemmV3WithAlphaBeta 直接路由至 GemmV3,并支持 broadcast bias 与 ND×NZ 输入组合。
  • 补齐 GemmV3 类型声明与 Kernel 配置:gemm_v3_def.cpp 新增 FP16/BF16 的 ND×ND 输入、FP32 输出及 ND×NZ 输入、低精度输出共四组类型/格式组合;ascend910b 和 ascend910_93 的 gemm_v3_binary.json 新增 GemmV3_ND_NZ_ND_ND_FP16_FP16_FP16_FP16 与 GemmV3_ND_NZ_ND_ND_BF16_BF16_BF16_BF16 两个同类型输出 Kernel。
  • 同步 tiling 结构字段命名:gemm_v3_tiling_data.h 与 gemm_v3_base_tiling.cpp 将 reservedBiasBroadcast 字段重命名为 reserved,并同步更新其赋值。
likedislike
不准确?
atomgit-bot
atomgit-bot
8月17日 评论:

代码审查

✅ 未发现问题

likedislike
不准确?
CANN-robotCANN-robot成员
8月17日 添加了label:stat/needs-squash
CANN-robotCANN-robot成员
8月17日 添加了label:cann-cla/yes
此处折叠了57条消息 查看更多
CANN-robot
CANN-robot成员
8月18日 评论:

The following users do not have permission to comment /lgtm or /approve on any module in this PR:
void_ptr

likedislike
CANN-robotCANN-robot成员
8月18日 添加了label:approved
CANN-robotCANN-robot成员
8月18日 关闭了关联的issue
CANN-robotCANN-robot成员
8月18日 合入了pull request
CANN-robot
CANN-robot成员
8月18日 评论:

Pull Request 已合并或已关闭。

If you want to solve this problem, you can click here to do it in the FAQs.

likedislike