已开启
Feat: 新增面向arch35的aclblasSgemmStridedBatched接口 #1
zihan0007创建于  7月20日
zihan0007成员
7月20日 创建

Background(背景信息)

算子功能

aclblasSgemmStridedBatched 为 Strided Batched GEMM(批处理矩阵乘)算子,属 BLAS Level 3(矩阵-矩阵操作)层级。对 batchCount 个规格一致(m/n/k/transA/transB 相同)的 FP32 矩阵乘执行批量计算:

C_i = alpha * op(A_i) * op(B_i) + beta * C_i,   i = 0, 1, ..., batchCount-1

其中:
  A_i = A + i * strideA   (元素偏移,列主序 Column-Major 存储)
  B_i = B + i * strideB
  C_i = C + i * strideC

  op(X) = X       当 trans = ACLBLAS_OP_N
  op(X) = X^T     当 trans = ACLBLAS_OP_T
  op(X) = X^H     当 trans = ACLBLAS_OP_C (实数 FP32 下 X^H == X^T)
  • batchCount:批数量,共执行 batchCount 个规格一致的 GEMM。
  • alphabeta:全批共用的 FP32 标量缩放因子。
  • strideA / strideB / strideC:相邻两个 batch 起始元素之间的偏移量(元素个数,非字节),支持 stride=0 广播语义(跨 batch 复用同一矩阵)。
  • 所有矩阵按列主序(Column-Major)存储,lda / ldb / ldc 为主维度,语义与 NETLIB/CBLAS BLAS 标准一致。

接口签名

aclblasStatus_t aclblasSgemmStridedBatched(
    aclblasHandle_t handle,
    aclblasOperation_t transA,
    aclblasOperation_t transB,
    int m,
    int n,
    int k,
    const float* alpha,
    const float* A,
    int lda,
    int64_t strideA,
    const float* B,
    int ldb,
    int64_t strideB,
    const float* beta,
    float* C,
    int ldc,
    int64_t strideC,
    int batchCount);

目标芯片与架构

项目 内容
目标芯片 Ascend 950(Ascend 950PR / Ascend 950DT)
目标架构 arch35(DAV_3510)
支持数据类型 FP32(A / B / C 均为 float,计算精度 FP32)
产品支持 Ascend 950PR / 950DT:支持;Atlas A2 / A3 系列:不支持

参数约束(摘要)

  • handle == nullptrACLBLAS_STATUS_HANDLE_IS_NULLPTR
  • transA / transB 不属于 {N, T, C} → ACLBLAS_STATUS_INVALID_VALUE
  • m < 0n < 0k < 0batchCount < 0ACLBLAS_STATUS_INVALID_VALUE
  • alpha == nullptrbeta == nullptrACLBLAS_STATUS_INVALID_VALUE
  • 主维度约束:lda >= max(1, transA==N ? m : k)ldb >= max(1, transB==N ? k : n)ldc >= max(1, m)
  • 边界快捷路径:m==0/n==0/batchCount==0 → 直接返回 SUCCESS;k==0alpha==0 → 仅缩放 C_i = beta * C_i

精度标准

  • 精度标准:浮点计算类社区标准,FP32 单标杆(CPU/NumPy FP32 golden)。
  • 精度指标阈值:
    • MARE(最大相对误差)≤ 10 × 2^-13 ≈ 1.221e-3
    • MERE(平均相对误差)≤ 2^-13 ≈ 1.221e-4
  • 近零抵消元素通过 atol_floor=2^-13 钳位豁免,避免近零元素相对误差固有过严。

Origin(信息来源)

cann 开发者

Benefit / Necessity(价值/作用)

Strided Batched GEMM 是介于单次 GEMM(aclblasGemmEx)与 Grouped Batched GEMM(aclblasSgemmGroupedBatched)之间的批处理形态——所有 batch 维度一致,仅通过固定 stride 定位,无需传入指针数组。应用场景包括:

  1. 批量矩阵乘:深度学习训练/推理中的多头注意力、批量线性变换等场景,多个规格一致的 GEMM 通过 stride 偏移紧凑排布于连续显存,一次调用完成全批计算。
  2. 广播语义strideA=0 / strideB=0 支持跨 batch 复用同一矩阵(如批量矩阵-向量乘、权重共享场景),减少显存占用与搬运。
  3. BLAS 标准完整性:补全 ops-blas 仓 Level 3 BLAS 批处理接口矩阵,与业界 cublasSgemmStridedBatched 接口惯例对齐(参数顺序:矩阵指针 + ld + stride 三元组 + batchCount 末置)。

Design(设计方案)

编程模型

采用 Blaze tensor_api 路径AscendC::Te CopyAtom/MmadAtom),架构核心验证为 DAV_3510,与目标 arch35 一致。GEMM 的数据搬运与矩阵乘全部由 tensor_api 实现,alpha/beta 后处理由独立 SIMD-membase Vector kernel 承接。

模板选型

  • 基线:纯AIC 模板——AIC 全流程 GM→L1→L0→MMAD→L0C→Fixpipe→GM,fast path(alpha==1,beta==0)可 Fixpipe 直写 C_i,零 AIV 开销。
  • StreamK:仅设计文档未交付可运行模板,不采用。
  • FixpOpti:登记为 general path 后续性能优化项,基线不采用。

Tiling 策略

  • batch 串行(Host 循环)+ 单 batch 内 Blaze GEMM 多核 M/N tile 切分。逐 batch 复用同一套 tiling 与同一 workspace。
  • 单 batch 内 serpentine(蛇形)二维 tile 切核(divM×divN 跨核轮转,奇数行 N 反向遍历复用 L1)。
  • tile 尺寸:tileM×tileN=128×128tileKChunk=256baseM/baseN=16baseK=8(FP32 C0=8)。
  • L1/L0 ping-pong 双缓冲,K-chunk 累加。

列主序处理

沿用 BLAS 列主序 swap 技巧:计算 C^T = op(B)^T op(A)^T,把「左矩阵=原 B、右矩阵=原 A」喂给行主序引擎,输出 C^T(n×m 行主序),物理内存等于 C(m×n 列主序)。引擎侧维度:mEff=原始 nnEff=原始 mkEff=原始 k

GM 输入采用 NDExt/DNExt layout pattern,按 transA/transB 四组合(NN/NT/TN/TT)分发,硬件自动完成 ND→NZ/ZN 格式转换,无需离线预重排。

alpha/beta 后处理(epilogue)

  • fast path(alpha==1 且 beta==0 且 ldc%8==0):Blaze GEMM Fixpipe 直写 C_i,无 workspace、无 combine kernel。
  • general path(alpha!=1 或 beta!=0):两段式——Blaze GEMM 写紧凑对齐 temp(workspace)→ 独立 SIMD-membase combine kernel 读 temp + 原始 C_i → alpha*temp+beta*C_i → 写回 C_i。
  • workspace 峰值:n * CeilAlign(m, 8) * 4 字节,逐 batch 复用同一 temp,不随 batchCount 增长。

关键设计点

  1. Fixpipe N 对齐约束(FIXPIPE_N_ALIGN=8):Fixpipe L0C→GM 对行 stride 有 8 元素(32B)对齐要求。fast path 直写用户 ldc 仅在 ldc%8==0 合法,否则退化对齐 temp+copy。
  2. 紧凑 temp 行 stride:列主序 swap 后引擎输出 C^T 列数 = nEff = 原始 m,故 temp 行 stride = CeilAlign(nEff, 8) = CeilAlign(原始 m, 8)(非 n),矩形 m≠n 场景关键。
  3. workspace 由 handle 统一管理:禁止 aclrtMalloc 额外分配,通过 aclblasGetEffectiveWorkspace(h) 获取,不足时返回错误码提示所需字节数。

参考实现

  • 仓内 blas/trmm/arch35/strmm_kernel.cpp(Blaze FP32 GEMM 已验证实现)为 GEMM kernel 直接改造基底。
  • blas/gemm/arch35/(列主序 + alpha/beta 后处理范式)、blas/gemm_grouped_batched/arch35/(batch 切分范式)为架构模式参考。
likedislike
Zzihan0007成员
7月20日 关联了pull request:Feat: 新增面向arch35的aclblasSgemmStridedBatched接口