Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
GRU(Gated Recurrent Unit)是常用的序列建模单元。当前 PyTorch 框架中 torch.nn.GRU 的 backward 包含 6 个输出(dx、dh_prev、dw_ih、dw_hh、db_ih、db_hh),现有实现通过逐个 kernel launch 组合计算,中间结果在 GM 间反复读写,访存带宽成为性能瓶颈。
torch.nn.GRU
本需求将 GRU 反向融合为单个 AscendC 算子 GruGrad,将 vector 门梯度计算 + 4 个 matmul(dgateMM、dwIhMM、dwHhMM、dxMM)+ bias reduce 整合到一颗 kernel 内,消除不必要的中间结果 GM 搬运,显著提升反向传播性能。
该需求来源于 PyTorch 大模型训练场景中的 GRU 反向性能优化。
1.通过 kernel 融合消除 5+ 次中间结果 GM 读写,访存带宽利用率有所提升。 2.支持不定长场景,补全完整性。
将 GRU 反向计算中的 4 类操作融合到单颗 AscendC kernel 内:
时间步循环(tIdx = T-1 → 0): 1. ProcessVector — 门梯度计算(d_reset/d_update/d_i_new/drh)+ 数据重排写回 2. ProcessDgateMM — d_gh × w_hh → dh_prev 的 matmul 部分 3. AccumulateDhPrev — dh_prev += grad_h * z 循环结束后: 4. ProcessDwIhMM — d_gi^T × x → dwInput 5. ProcessDwHhMM — d_gh^T × h_prev → dwHidden 6. ProcessDxMM — d_gi × w_ih^T → dx 7. ProcessBiasReduce × 2 — sum(d_gi) / sum(d_gh) → dbInput / dbHidden
Process(): // 预配置: dgateMM 的 B 矩阵(w_hh)和 tail 不随时间步变化 if GetBlockIdx() < dgateMMTiling.usedCoreNum: dgateMM.SetTensorB(wHidden, false) if not isSeqLength: ApplyTail(dgateMM, dgateMMTiling, dgateTail) InitDhPrev() // 不定长: 预写 dh 到 dhPrevWsGm for tIdx = T-1 down to 0: ProcessVector(tIdx) -- 门梯度计算 + 数据重排 SyncAll() ProcessDgateMM(tIdx) -- d_gh * w_hh -> dhPrevWs SyncAll() AccumulateDhPrev(tIdx)-- dhPrevWs += dhFromH StoreDhPrev() -- dhPrevWs -> dhPrev 输出 ProcessDwIhMM() -- d_gi^T * x -> dwInput ProcessDwHhMM() -- d_gh^T * h -> dwHidden ProcessDxMM() -- d_gi * w_ih^T -> dx if isBias == 1: ProcessBiasReduce(dGiGm, dbInput, rows=totalSteps, cols=3H) ProcessBiasReduce(dGhGm, dbHidden, rows=totalSteps, cols=3H)
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
GRU(Gated Recurrent Unit)是常用的序列建模单元。当前 PyTorch 框架中
torch.nn.GRU的 backward 包含 6 个输出(dx、dh_prev、dw_ih、dw_hh、db_ih、db_hh),现有实现通过逐个 kernel launch 组合计算,中间结果在 GM 间反复读写,访存带宽成为性能瓶颈。本需求将 GRU 反向融合为单个 AscendC 算子 GruGrad,将 vector 门梯度计算 + 4 个 matmul(dgateMM、dwIhMM、dwHhMM、dxMM)+ bias reduce 整合到一颗 kernel 内,消除不必要的中间结果 GM 搬运,显著提升反向传播性能。
Origin(信息来源)
该需求来源于 PyTorch 大模型训练场景中的 GRU 反向性能优化。
Benefit / Necessity (价值/作用)
1.通过 kernel 融合消除 5+ 次中间结果 GM 读写,访存带宽利用率有所提升。
2.支持不定长场景,补全完整性。
Design(设计方案)
将 GRU 反向计算中的 4 类操作融合到单颗 AscendC kernel 内:
Process 整体流程