已关闭
[Requirement|需求建议]: Addmm、Baddbmm 和 AddmmWeightNz支持broadcast bias输入走入gemmv3 kernel #4844
HKFLYE创建于 20 天前关闭于 20 天前
20 天前 添加了label:requirement
20 天前 将 HuangKun8682 设为负责人
20 天前 修改标题为 “[Requirement|需求建议]: Addmm、Baddbmm 和 AddmmWeightNz支持broadcast bias输入走入gemmv3 kernel”,原标题为“[Requirement|需求建议]: ”
20 天前 修改标题为 “[Requirement|需求建议]: Addmm、Baddbmm 和 AddmmWeightNz支持broadcast bias输入走入gemmv3 kernel”,原标题为“[Requirement|需求建议]: ”
20 天前 关联了pull request:【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3
20 天前 关联了pull request:【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3
20 天前 关闭了 issue
19 天前 添加了label:resolved
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
当前 A2/A3 上 Addmm、Baddbmm 和 AddmmWeightNz 在 bias shape 与矩阵乘输出 shape 不一致时,部分场景无法直接使用 GemmV3:
broadcast bias 会回退到 MatMul/BatchMatMul、Muls、Axpy 等小算子拼接通路,增加 Kernel Launch 和中间结果读写开销。
AddmmWeightNz 在 cubeMathType=USE_FP32_ADD 时会被转换为 KEEP_DTYPE,无法使用 GemmV3 的 FP32 中间累加能力。
GemmV3 缉少部分 ND×ND、ND×NZ 类型声明及 A2/A3 binary 配置,导致接口即使完成路由,也可能无法找到对应 Kernel。
GemmV3 Kernel 原有 bias 处理要求 bias shape 与输出完全一致,不能在 Kernel 内完成 N、M、scalar 和 per-batch scalar broadcast。
新需求是在 A2/A3(DAV_2201)上扩展 GemmV3 对 broadcast bias 的支持,并完成 Addmm、Baddbmm、AddmmWeightNz 等 ACLNN 接口适配,使矩阵乘、alpha/beta 缩放和 bias broadcast 能够在一个 GemmV3 Kernel 中完成。
Origin(信息来源)
用户诉求
Benefit / Necessity (价值/作用)
将以下计算:
PPMatMul/BatchMatMul + Muls + Axpy
融合为:
GemmV3(A, B, bias, alpha, beta)
可以减少:
Kernel Launch 次数;
矩阵乘中间结果的GM落盘和重新读取;
Muls、Axpy 等小算子的调度开销;
broadcast bias 的额外展开和数据搬运。
Design(设计方案)