已关闭
Feat: 新增面向ascend950的aclblasSgemmEx接口 #308
wangzitao创建于  7月13日关闭于  7月15日
wangzitao
wangzitao成员
7月13日 创建

Background(背景信息)

新增 aclblasSgemmEx 算子,实现单精度(FP32)通用矩阵乘法,核心运算为:

C = alpha * op(A) * op(B) + beta * C

其中矩阵 A、B、C 以及标量 alpha、beta 均为 FP32 类型,采用 BLAS 标准列主序存储。Ex 后缀表示保留算法选择参数 algo,与仓内已有 aclblasGemmEx 接口对齐,当前仅支持 ACLBLAS_GEMM_DEFAULT。S 前缀表示单精度(FP32)专用,A/B/C/alpha/beta 均为 float 类型,无需传入 Atype/Btype/Ctype/computeType 参数。

接口签名

aclblasStatus_t aclblasSgemmEx(
    aclblasHandle_t handle,
    aclblasOperation_t transA, aclblasOperation_t transB,
    int m, int n, int k,
    const float* alpha,
    const float* A, int lda,
    const float* B, int ldb,
    const float* beta,
    float* C, int ldc,
    aclblasGemmAlgo_t algo);

参数规格

  • transA/transB:矩阵操作类型,支持 ACLBLAS_OP_N(不转置)、ACLBLAS_OP_T(转置)、ACLBLAS_OP_C(共轭转置,FP32 实数等价于转置)
  • m/n/k:矩阵维度,M ≥ 0、N ≥ 0、K ≥ 0
  • alpha/beta:FP32 标量系数指针,不可为 nullptr
  • A/B/C:FP32 设备内存指针,列主序存储
  • lda/ldb/ldc:矩阵主维度(列主序),需满足 BLAS 标准约束
  • algo:算法选择,当前仅支持 ACLBLAS_GEMM_DEFAULT

精度标准

  • MARE ≤ 10 × 2^-13(最大绝对相对误差)
  • MERE ≤ 2^-13(最大相对误差)

目标芯片与 CANN 版本

项目 内容
目标芯片 Ascend950(Ascend950PR / Ascend950DT)
目标架构 arch35(DAV_3510)
CANN 版本 9.1.0

Origin(信息来源)

cann 开发者

Benefit / Necessity(价值/作用)

BLAS(Basic Linear Algebra Subprograms)Level 3 标准矩阵乘法是科学计算、深度学习训练与推理中最基础且高频调用的运算原语之一。aclblasSgemmEx 提供 FP32 单精度矩阵乘法接口,支持矩阵转置组合(N/T/C)和 alpha/beta 标量缩放融合,满足以下应用场景需求:

  1. 科学计算:线性代数求解、最小二乘法、特征值分解等依赖 GEMM 作为核心计算单元
  2. 深度学习:全连接层、注意力机制中的 Q×K^T 和 attn×V 计算、矩阵乘法融合算子的基础实现
  3. 工程应用:信号处理中的 FFT 后处理、图像处理中的仿射变换等

该算子针对 Ascend950(arch35)架构进行适配,利用 Cube MMA 阵列实现硬件加速矩阵乘法,同时通过 Vector Kernel 融合 alpha/beta 后处理,减少额外 Kernel launch 开销和中间结果访存。接口遵循 BLAS 标准列主序约定,便于既有 BLAS 应用程序迁移至昇腾平台。

Design(设计方案)

编程模型

SIMD membase(传统 SIMD,TPipe/TQue 流水线 + BlockMmad 低阶 API),采用双 Kernel 架构:

Kernel 类型 职责
sgemm_ex_cube_kernel Cube Kernel 执行矩阵乘 op(B')×op(A'),结果写 GM
sgemm_ex_alpha_beta_kernel Vector Kernel 执行 C = alpha*tempAB + beta*C 后处理

列主序适配(Column-Major Trick)

BLAS 接口为列主序,NPU Cube 硬件按行主序分块计算。利用矩阵转置等价关系,Host 侧在 Tiling 计算后执行参数交换:swap(m, n)swap(lda, ldb)swap(isTransA, isTransB)、设备指针 aDevicePtr = BbDevicePtr = A。交换后 Cube Kernel 按行主序计算,结果写入 GM 即为列主序的 C。

Tiling 策略

  • 多核切分:M×N 输出矩阵按 2D 分块(mBlocks × nBlocks)分配到多 Cube 核
  • FP32 分块参数:baseM=32, baseN=16, baseK=8,与 Cube MMA 阵列 16×16×16 分形匹配
  • L0C 2D 分块:控制 L0C 占用 ≤ 256KB
  • K 维核内循环:K 维不切分核间,核内按 baseK 循环累加

分支场景覆盖

  • 正常 GEMM(alpha=1.0, beta=0.0):Cube Kernel 直接写回 C,无需后处理
  • alpha/beta 后处理:Cube Kernel 写 tempAB 到 workspace,Vector Kernel 执行后处理
  • 边界情况:m==0/n==0 直接返回;k==0/alpha==0 跳过矩阵乘执行 C=beta*C
  • 转置组合:transA/transB 的 N/T/C 全组合
likedislike
wangzitaowangzitao成员
7月13日 关联了pull request:Feat: 新增面向arch35的aclblasSgemmEx接口
CANN-robotCANN-robot成员
7月15日 关闭了 issue
CANN-robotCANN-robot成员
7月15日 添加了label:resolved