Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
BatchMatMul V3 算子新增 ITER_BATCH_BROADCAST 模板,解决BMMV3算子在batch broadcast场景无法做IterBatch(多batch聚合搬入/搬出)的问题,优化BMMV3算子在batch broadcast场景的性能
昇腾算子开发团队
现有 BatchMatMul V3 算子提供两条独立路径:
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))
/assign @justsozl
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 算子提供两条独立路径:
当输入同时存在广播和大量 batch 时(如 A batch=1,64、B batch=8,64、C batch=8,64),两条路径都无法高效处理——BROADCAST_BATCH 逐批次搬运浪费 Cube 算力,ITER_BATCH 直接拒绝广播场景。
需求价值
应用场景
大语言模型(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))