当前 ops-blas 库已实现面向 arch35 (Ascend950) 平台的单批次及批量 GEMM 接口,但缺少 分组批量(Grouped Batched) 版本。在实际应用中(如多专家路由、异构 batch 推理等场景),用户需要对多个分组内的矩阵批次执行 GEMM,且不同分组可拥有不同的维度、转置方式及缩放因子。若缺少分组接口,用户只能在 Host 侧按组循环调用单批次或同参批量接口,引入额外的 kernel launch 开销,且无法在一次调用中统一调度所有分组任务。
aclblasSgemmGroupedBatched 是 FP32 分组批量矩阵乘接口,支持通过设备侧指针数组指定多组矩阵,在单次调用中完成所有分组内批次的 GEMM 计算:
aclblasSgemmGroupedBatched
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) 接口。
cann 开发者
cublasSgemmGroupedBatched
总体架构:Host 侧负责参数校验、Tiling 计算与元数据上传;Kernel 侧采用 AIV SIMD membase 模式完成全部数值计算。
Kernel 计算路径:
多核切分:按总 batch 数(sum(groupSize[g]))均匀分配到各 Vector Core,grid-stride 循环处理。
测试框架扩展:
csv_loader.h
fill.h
Backgroud(背景信息)
当前 ops-blas 库已实现面向 arch35 (Ascend950) 平台的单批次及批量 GEMM 接口,但缺少 分组批量(Grouped Batched) 版本。在实际应用中(如多专家路由、异构 batch 推理等场景),用户需要对多个分组内的矩阵批次执行 GEMM,且不同分组可拥有不同的维度、转置方式及缩放因子。若缺少分组接口,用户只能在 Host 侧按组循环调用单批次或同参批量接口,引入额外的 kernel launch 开销,且无法在一次调用中统一调度所有分组任务。
aclblasSgemmGroupedBatched是 FP32 分组批量矩阵乘接口,支持通过设备侧指针数组指定多组矩阵,在单次调用中完成所有分组内批次的 GEMM 计算:每个分组 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 (价值/作用)
cublasSgemmGroupedBatched等业界接口语义对齐Design(设计方案)
总体架构:Host 侧负责参数校验、Tiling 计算与元数据上传;Kernel 侧采用 AIV SIMD membase 模式完成全部数值计算。
Kernel 计算路径:
多核切分:按总 batch 数(sum(groupSize[g]))均匀分配到各 Vector Core,grid-stride 循环处理。
测试框架扩展:
csv_loader.h:新增分号分隔数组解析(parseIntArray、parseFloatArray、parseTransArray 等),支持 grouped 参数 CSV 驱动fill.h:保留 FP8/FP4/INT 专用填充,扩展 Double/Complex 填充函数