aclblasStatus_t aclblasSgemmStridedBatched(
aclblasHandle_t handle,
aclblasOperation_t transA,
aclblasOperation_t transB,
int m,
int n,
int k,
constfloat* alpha,
constfloat* A,
int lda,
int64_t strideA,
constfloat* B,
int ldb,
int64_t strideB,
constfloat* beta,
float* C,
int ldc,
int64_t strideC,
int batchCount);
Background(背景信息)
算子功能
aclblasSgemmStridedBatched为 Strided Batched GEMM(批处理矩阵乘)算子,属 BLAS Level 3(矩阵-矩阵操作)层级。对batchCount个规格一致(m/n/k/transA/transB 相同)的 FP32 矩阵乘执行批量计算:batchCount:批数量,共执行batchCount个规格一致的 GEMM。alpha、beta:全批共用的 FP32 标量缩放因子。strideA / strideB / strideC:相邻两个 batch 起始元素之间的偏移量(元素个数,非字节),支持stride=0广播语义(跨 batch 复用同一矩阵)。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);目标芯片与架构
参数约束(摘要)
handle == nullptr→ACLBLAS_STATUS_HANDLE_IS_NULLPTRtransA / transB不属于 {N, T, C} →ACLBLAS_STATUS_INVALID_VALUEm < 0或n < 0或k < 0或batchCount < 0→ACLBLAS_STATUS_INVALID_VALUEalpha == nullptr或beta == nullptr→ACLBLAS_STATUS_INVALID_VALUElda >= 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==0或alpha==0→ 仅缩放C_i = beta * C_i精度标准
atol_floor=2^-13钳位豁免,避免近零元素相对误差固有过严。Origin(信息来源)
cann 开发者
Benefit / Necessity(价值/作用)
Strided Batched GEMM 是介于单次 GEMM(
aclblasGemmEx)与 Grouped Batched GEMM(aclblasSgemmGroupedBatched)之间的批处理形态——所有 batch 维度一致,仅通过固定 stride 定位,无需传入指针数组。应用场景包括:strideA=0/strideB=0支持跨 batch 复用同一矩阵(如批量矩阵-向量乘、权重共享场景),减少显存占用与搬运。cublasSgemmStridedBatched接口惯例对齐(参数顺序:矩阵指针 + ld + stride 三元组 + batchCount 末置)。Design(设计方案)
编程模型
采用 Blaze tensor_api 路径(
AscendC::TeCopyAtom/MmadAtom),架构核心验证为 DAV_3510,与目标 arch35 一致。GEMM 的数据搬运与矩阵乘全部由 tensor_api 实现,alpha/beta 后处理由独立 SIMD-membase Vector kernel 承接。模板选型
Tiling 策略
divM×divN跨核轮转,奇数行 N 反向遍历复用 L1)。tileM×tileN=128×128、tileKChunk=256、baseM/baseN=16、baseK=8(FP32 C0=8)。列主序处理
沿用 BLAS 列主序 swap 技巧:计算
C^T = op(B)^T op(A)^T,把「左矩阵=原 B、右矩阵=原 A」喂给行主序引擎,输出 C^T(n×m 行主序),物理内存等于 C(m×n 列主序)。引擎侧维度:mEff=原始 n、nEff=原始 m、kEff=原始 k。GM 输入采用 NDExt/DNExt layout pattern,按 transA/transB 四组合(NN/NT/TN/TT)分发,硬件自动完成 ND→NZ/ZN 格式转换,无需离线预重排。
alpha/beta 后处理(epilogue)
alpha*temp+beta*C_i→ 写回 C_i。n * CeilAlign(m, 8) * 4字节,逐 batch 复用同一 temp,不随 batchCount 增长。关键设计点
ldc%8==0合法,否则退化对齐 temp+copy。nEff= 原始 m,故 temp 行 stride =CeilAlign(nEff, 8) = CeilAlign(原始 m, 8)(非 n),矩形 m≠n 场景关键。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 切分范式)为架构模式参考。