已关闭
[Requirement|需求建议]: 新增 STRSMBATCHED (FP32) 算子 DAV_3510 架构实现 #335
陶贻坤创建于  7月28日关闭于  7月30日
陶贻坤
陶贻坤成员
7月28日 创建

Backgroud(背景信息)

BLAS Level-3 批量三角方程组求解接口 aclblasStrsmbatched 当前缺少 DAV_3510(ascend950)架构的支持。本次新增 arch35 实现,使 Strsmbatched 在 ascend950 上可用。

Strsmbatched 计算公式:

  • side=LEFT: op(A_i) * X_i = alpha * B_i
  • side=RIGHT: X_i * op(A_i) = alpha * B_i

其中 A_i 为 batchCount 个三角矩阵(仅存储 UPPER 或 LOWER 一个三角),B_i 为普通矩阵,op(A) 可为 A / A^T / A^H。

实现采用与已有 aclblasStrsm 一致的 SIMT + tensor_api 混合架构:

  1. Panel Solve(AIV核, SIMT):对 panel 块进行三角前代/回代求解
  2. Trailing Update GEMM(AIC核, tensor_api):用 Cube GEMM 计算 trailing matrix 更新
  3. Scale/Transpose/Extract/AXPY(AIV核, SIMT):辅助 kernel

涉及新增文件:

  • blas/trsmbatched/arch35/strsmbatched_host.cpp — Host 侧入口(参数校验、Tiling 计算、路径选择、kernel 调度)
  • blas/trsmbatched/arch35/strsmbatched_kernel.cpp — Kernel 侧实现(panel/scale/zero/extract/axpy/transpose/gemm 共 8 个 kernel)
  • blas/trsmbatched/arch35/strsmbatched_kernel.h — kernel 声明
  • blas/trsmbatched/arch35/strsmbatched_tiling_data.h — TilingData 结构体定义
  • blas/trsmbatched/README.md — 算子文档
  • test/trsmbatched/strsmbatched/arch35/strsmbatched_test.cpp — CSV 参数化精度 ST
  • test/trsmbatched/strsmbatched/arch35/strsmbatched_npu_wrapper.h — H2D/D2H 测试 wrapper
  • test/trsmbatched/strsmbatched/arch35/strsmbatched_test.csv — 58 条测试用例(含边界值和非法参数校验)
  • test/trsmbatched/strsmbatched/strsmbatched_golden.h — golden 计算参考
  • test/trsmbatched/strsmbatched/strsmbatched_param.h — 测试参数定义
  • test/trsmbatched/strsmbatched/CMakeLists.txt — 构建配置

Origin(信息来源)

cann开发者

Benefit / Necessity (价值/作用)

  1. 填补 aclblasStrsmbatched 在 ascend950 (DAV_3510) 上的缺失,完善 BLAS Level-3 batched 接口覆盖率
  2. SIMT + tensor_api 混合架构充分利用 ascend950 Cube+Vector 双硬件能力,与已有 aclblasStrsm 路线保持一致
  3. Blocked path 采用 panel SIMT + Cube GEMM trailing update,GEMM kernel 使用 L0A/L0B ping-pong + L1 双缓冲流水线
  4. 支持 side=LEFT/RIGHT、uplo=UPPER/LOWER、trans=N/T/C、diag=UNIT/NON_UNIT 全组合,alpha 支持零值快速路径
  5. Right 路径通过 transpose 归约为 Left 求解,复用 Left 的 SIMT/Blocked 两条子路径
  6. 58 条 ST 覆盖基本功能、边界值(m=0/n=0)、side/uplo/trans/diag 全组合、alpha=0 快速路径等场景

Design(设计方案)

路径选择策略

side 条件 路径
LEFT m > 128 或 (m <= 128 && n >= 256) Blocked(panel SIMT + Cube GEMM trailing update)
LEFT 否则 SIMT-only(panel kernel 全量求解)
RIGHT n <= 128 RightSimt(transpose → left simt solve → transpose back)
RIGHT 否则 RightDevice(transpose → left blocked solve → transpose back)
* alpha == 0 zero kernel 快速清零 B

Panel Solve Kernel (AIV-only, SIMT)

按列分配到 AIV 核,每核处理 colStart..colEnd 列:

  • 将 panel 块的 A 三角部分加载到 UB(__ubuf__ float aBlockUb[])
  • 按行迭代求解,根据 uplo/trans 确定 forward/backward 遍历方向
  • 每行计算 dot product 减去已知解的贡献,再除以对角元(diag=NON_UNIT 时)
  • 对角元为 1(diag=UNIT)时跳除法

Trailing Update GEMM Kernel (AIC-only, tensor_api)

四层循环结构:核级 tile → MN tile → K chunk → L0 K split

  • 多核切分:ceilDiv(m, tileM) × ceilDiv(n, tileN) 个 tile,各核步进式领取,Snake Schedule 奇数行反转 N 方向
  • L1 双缓冲:A-side 和 B-side 数据按 NZ/ZN 格式写入共享 L1 buffer,双 buffer ping-pong
  • L0 LoadData → MMAD:K chunk 进一步按 BASE_K=8 切分到 L0 级,L0A/L0B ping-pong,unitFlag 控制累加/清零
  • L0C NZ2ND 输出:CopyL0C2GM 将 NZ 格式转为 row-major 写入 temp GM
  • 四级流水线同步:MTE2(搬入L1) → MTE1(搬入L0) → M(MMAD) → FIX(搬出GM),SetFlag/WaitFlag 配对

NoTrans vs Trans Trailing Update

模式 GEMM 输入 数据流
NoTrans B(k×mC) × A(bs×k) → temp GEMM 直接从 GM 读取 B/A,结果 AXPY 回 B
Trans 先 extract A/B 到连续 workspace,再 GEMM(mC×n) × (bs×n) → temp,最后 AXPY 回 B

Right 路径(Transpose 归约)

Right TRSM: X * op(A) = alpha*B → 转置 → op(A)^T * X^T = alpha*B^T

  1. Scale B by alpha
  2. Transpose B → Bt
  3. Left solve: op(A)^T * Bt = Bt(flip trans: N↔T,uplo 不变)
  4. Transpose Bt → B

Host 侧

  • 参数校验(side/uplo/trans/diag/m/n/alpha/Aarray/Barray/lda/ldb/batchCount)
  • Aarray/Barray 从 device 拷贝到 host(aclrtMemcpy D2H),逐 batch 调用 LaunchTrsmbatchedKernel
  • 动态核数分配(min(数据量, 硬件核数))
  • workspace 按路径动态分配(temp/aWs/bWs/Bt)
likedislike
陶贻坤陶贻坤成员
7月28日 将 eternityk 设为负责人
陶贻坤陶贻坤成员
7月28日 修改了issue 的描述
CANN-robotCANN-robot成员
7月30日 关闭了 issue
CANN-robotCANN-robot成员
7月30日 添加了label:resolved