已关闭
[Requirement|需求建议]: Feat: 新增面向arch35的aclblasSgemmGroupedBatched接口 #187
Crrryyyy创建于  6月22日关闭于  6月24日
Crrryyyy
Crrryyyy成员
6月22日 创建

Backgroud(背景信息)

当前 ops-blas 库已实现面向 arch35 (Ascend950) 平台的单批次及批量 GEMM 接口,但缺少 分组批量(Grouped Batched) 版本。在实际应用中(如多专家路由、异构 batch 推理等场景),用户需要对多个分组内的矩阵批次执行 GEMM,且不同分组可拥有不同的维度、转置方式及缩放因子。若缺少分组接口,用户只能在 Host 侧按组循环调用单批次或同参批量接口,引入额外的 kernel launch 开销,且无法在一次调用中统一调度所有分组任务。

aclblasSgemmGroupedBatched 是 FP32 分组批量矩阵乘接口,支持通过设备侧指针数组指定多组矩阵,在单次调用中完成所有分组内批次的 GEMM 计算:

C[i] = alpha[g] × op(A[i]) × op(B[i]) + beta[g] × C[i],  i 属于分组 g

每个分组 g 可独立配置 (m, n, k, transa, transb, alpha, beta, lda, ldb, ldc) 及 groupSize[g];同一分组内各 batch 共享该组参数。所有矩阵按 Column-Major(列主序)存储。

目标平台为 Ascend950 (arch35 / DAV_3510),当前仅保留 S (FP32) 接口。

Origin(信息来源)

cann 开发者

Benefit / Necessity (价值/作用)

  1. 分组调度效率提升:单次 kernel launch 完成多组、多 batch 的 GEMM,避免 Host 侧按组循环调用带来的 launch 开销
  2. 灵活的参数配置:不同分组可使用不同的 (m, n, k)、转置标志和 alpha/beta,满足异构 batch 场景
  3. 接口完整性:补齐 BLAS Level-3 分组批量 GEMM 接口,与 cuBLAS cublasSgemmGroupedBatched 等业界接口语义对齐
  4. 多核负载均衡:总 batch 数均匀分配到各 Vector Core,自动适应单组多 batch 与多组少 batch 等场景

Design(设计方案)

总体架构:Host 侧负责参数校验、Tiling 计算与元数据上传;Kernel 侧采用 AIV SIMD membase 模式完成全部数值计算。

Kernel 计算路径

  • transa=N:列逐步累加(ColumnWise),逐列读取 A/B tile 并向量乘加写入 C
  • transa=T:行点积(DotProduct),逐行向量点积后 ReduceSum 写入 C
  • TT 路径:内部交换 A↔B、m↔n,等效为 NN 模式处理
  • alpha=0 或 k=0:走 BetaOnly 路径,仅执行 C = beta × C

多核切分:按总 batch 数(sum(groupSize[g]))均匀分配到各 Vector Core,grid-stride 循环处理。

测试框架扩展

  • csv_loader.h:新增分号分隔数组解析(parseIntArray、parseFloatArray、parseTransArray 等),支持 grouped 参数 CSV 驱动
  • fill.h:保留 FP8/FP4/INT 专用填充,扩展 Double/Complex 填充函数
likedislike
wangzitaowangzitao成员
6月22日 将 Crrryyyy 设为负责人
CANN-robotCANN-robot成员
6月24日 关闭了 issue
CANN-robotCANN-robot成员
6月24日 添加了label:resolved