已关闭
[Requirement|需求建议]: BMMV3在batch broadcast场景支持多batch搬入搬出 #3849
justsozl创建于  7月4日关闭于  7月7日
justsozl成员
7月4日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

BatchMatMul V3 算子新增 ITER_BATCH_BROADCAST 模板,解决BMMV3算子在batch broadcast场景无法做IterBatch(多batch聚合搬入/搬出)的问题,优化BMMV3算子在batch broadcast场景的性能

Origin(信息来源)

昇腾算子开发团队

Benefit / Necessity (价值/作用)

现有 BatchMatMul V3 算子提供两条独立路径:

  • BROADCAST_BATCH:处理广播场景(A/B 单侧 batch 维度用 % 映射),但逐批次搬运,无 L1/L0 流水线聚合
  • ITER_BATCH:处理 L1/L0 多批次聚合(提升 Cube 利用率),但要求 A/B batch 维度完全一致,不支持广播
    当输入同时存在广播和大量 batch 时(如 A batch=1,64、B batch=8,64、C batch=8,64),两条路径都无法高效处理——BROADCAST_BATCH 逐批次搬运浪费 Cube 算力,ITER_BATCH 直接拒绝广播场景。
    需求价值
  1. 性能提升
  • 广播侧数据在 L1 中只存一份(LAST_BATCH_DIM 固定偏移),非广播侧可加载更多批次到 L1,减少 GM→L1 搬运次数
  • L0 多批次聚合(iterBatchL0)使 Cube 在多批次上流水线执行,减少单批次 MMAD→fixpipe 的空闲间隔
  • 实测场景下(如 A 广播、batch 维度较大),相比 BROADCAST_BATCH 逐批处理,吞吐量提升显著
  1. 覆盖关键缺失场景
  • 补全了"广播+大 batch"的空白组合,此前该场景只能走低效的逐批路径或直接 fallback
  1. 内存效率
  • 广播侧共享 L1 数据,同等 L1 容量下可容纳更多非广播侧批次,提高 iterBatchL1 上限
  • L1 ping-pong 双缓冲下广播数据固定存放,避免半个 L1 空间浪费
    应用场景
    大语言模型(LLM)推理 / 推荐系统

Design(设计方案)

Kernel 入口
for tileIdx = blockIdx; tileIdx < tileNum; tileIdx += blockNum:
aGmStart = aBroadcast ? ComputeABroadcastIndex(startBatch) : startBatch
bGmStart = bBroadcast ? ComputeBBroadcastIndex(startBatch) : startBatch
gmASlice = Slice(gmA, aGmStart, al1Count) // 广播侧可能只 Slice 1 批
gmBSlice = Slice(gmB, bGmStart, bl1Count)
gmCSlice = Slice(gmC, startBatch, curIterBatchL1)
blockMmad(gmCSlice, gmASlice, gmBSlice, gmBias, curIterBatchL1)
MMAD 入口
l1BufId = ping-pong id (0/1)
offsetAl1 = HALF_L1 * l1BufId; offsetBl1 = offsetAl1 + aL1Size
Copy GM→L1: A/B/bias (MTE2 同步)
Copy L12L0 objects, mmadAtom 初始化
if iterBatchL0 > 1 → BatchedMmadLoop
else → SingleBatchMmadLoop
Batched 模式 (iterBatchL0 > 1, 无 MNK 切分)
for iter1 = 0; iter1 < CeilDiv(curIterBatchL1, iterBatchL0); iter1++:
l0BufId = iter1 % 2
Copy L1→L0 (3D or 2D per-batch, 广播侧 A/B 只有1批时走 2D)
for batchL0Idx in curIterBatchL0:
MMAD(M, N, K, L0A_off[batchL0Idx], L0B_off[batchL0Idx], L0C_off[batchL0Idx])
FixL0CToGM (3D, 整批写回 gmC)
Single-batch 模式 (iterBatchL0 == 1, MNK 切分)
al1Off = (broadcastAxisA==3) ? 0 : offsetAl1 // 广播侧固定偏移
bl1Off = (broadcastAxisB==3) ? aL1Size : al1Off + aL1Size
for batchIdx; for iterN; for iterM:
for iterK:
l0BufId = (l0cBufIdx * kl0Cnt + iterK) % 2 // 全局轮转避免断流
Copy L1→L0 (2D per-batch, 字节偏移计算)
MMAD(curM, curN, curK, iterK) // K 累加
FixL0CToGMSingleBatch (2D, gmDest shape=(m,n))

likedislike
Jjustsozl成员
7月4日 关联了pull request:support bmm iterbatch broadcast
yuning_chen
yuning_chen成员
7月5日 评论:

/assign @justsozl

likedislike
CANN-robotCANN-robot成员
7月5日 将 justsozl 设为负责人
Jjustsozl成员
7月6日 关联了pull request:support bmm iterbatch broadcast
Jjustsozl成员
7月7日 issue状态由 进行中 改变为 已完成
Jjustsozl成员
7月7日 关闭了 issue
CANN-robotCANN-robot成员
7月7日 添加了label:Accepted