ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

CANN ops-transformer 中的 BlockAttentionResidualsGrad:融合 Softmax 与 RMSNorm 反向的注意力残差梯度算子解析

CANN ops-transformer 中的 BlockAttentionResidualsGrad:融合 Softmax 与 RMSNorm 反向的注意力残差梯度算子解析 CANN ops-transformer 中的 BlockAttentionResidualsGrad融合 Softmax 与 RMSNorm 反向的注意力残差梯度算子解析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 仓库中mhc/block_attention_residuals_grad算子目录展开系统讲解反向算子BlockAttentionResidualsGrad的数学原理、参数语义、约束条件以及基于 aclnn 两段式接口和 PyTorch 扩展的完整调用方法并深入到 op_def 注册、Host Tiling、Kernel 模板分派与 Workspace 分配等源码实现细节。读完本文你将掌握在 Ascend NPU 上为注意力残差Block Attention Residuals前向算子编写、调用与调试反向梯度计算的完整实战方案。算子定位注意力残差前向的反向梯度计算BlockAttentionResidualsGrad是 CANN ops-transformer 项目中正向算子 BlockAttentionResiduals注意力残差的反向传播算子。正向前向算子将partialBlock与blockRes按 block 维拼接成 value 序列后依次完成 RMS 归一化、normWeight与projWeight逐元素乘积投影打分、Softmax 归一化以及加权求和融合输出hiddenStates当needBackward为 true 时还会输出反向所需的两份中间量invNorm前向逐行 RMS 归一化系数shape 为[T, N1]仅 FLOAT32probs前向 Softmax 输出概率shape 为[T, N1]仅 FLOAT32。反向算子BlockAttentionResidualsGrad正是利用这两份前向保存的中间量结合上游梯度gradHiddenStates与权重projWeight、normWeight一次性完成 Softmax 反向、RMS 归一化反向与注意力加权求和反向的融合计算输出四个梯度张量gradPartialBlock、gradBlockRes、gradProjWeight、gradNormWeight从源码结构看该算子的完整实现分布如下目录/文件职责docs/aclnnBlockAttentionResidualsGrad.mdaclnn 接口级说明函数原型、参数表、返回码与调用示例examples/两个可直接参考的 C 调用样例普通路径与 SPLIT_H 路径op_host/算子定义注册、InferShape、Host Tiling 与 aclnn 接口实现op_kernel/AscendC Kernel 入口与 arch22/arch35 两套内核实现torch_extension/PyTorch 侧封装block_attention_residuals_backward接口产品支持情况根据算子目录 README.md 与 aclnn 接口文档 aclnnBlockAttentionResidualsGrad.md 中的产品支持矩阵产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×在算子定义注册文件 block_attention_residuals_grad_def.cpp 中可以看到AICore 配置分别注册了ascend910b对应 Atlas A2 训练/推理系列、ascend910_93对应 Atlas A3 系列与ascend950三条路径其中ascend950使用独立的 regbase Kernel 入口文件block_attention_residuals_grad_apt其余平台使用统一的 Kernel 入口文件block_attention_residuals_grad。这从源码层面印证了上表的产品支持矩阵。数学原理反向计算公式反向计算中首先将前向拼接 value 与归一化中间量重新定义如下$T$ 为 token 数$N$ 为blockRes的 block 数$H$ 为 hidden size$N1$ 个 value 由 $N$ 个分块残差加 1 个前缀和组成$$ V_{t,i,h} \begin{cases} block_res_{t,i,h}, i N \ partial_block_{t,h}, i N \end{cases} $$$$ k_{t,i,h} V_{t,i,h} \cdot inv_norm_{t,i} $$$$ score_weight_{h} norm_weight_{h} \cdot proj_weight_{0,h} $$第 1 步out对V、probs的反向梯度$$ g_{t,i} \sum_{h0}^{H-1} grad_output_{t,h} \cdot V_{t,i,h} $$$$ grad_score_{t,i} probs_{t,i} \cdot \left(g_{t,i} - \sum_{j0}^{N} probs_{t,j} \cdot g_{t,j}\right) $$其中 $g_{t,i}$ 是加权和 $out_{t,h}\sum_i probs_{t,i} V_{t,i,h}$ 对 $V$ 的梯度grad_score是标准 Softmax 交叉熵反向公式$p_i(g_i - \sum_j p_j g_j)$。第 2 步经score_weight与 RMS 归一化回传的梯度$$ grad_k_{t,i,h} grad_score_{t,i} \cdot score_weight_{h} $$$$ grad_score_weight_{h} \sum_{t0}^{T-1}\sum_{i0}^{N} grad_score_{t,i} \cdot k_{t,i,h} $$$$ grad_inv_norm_{t,i} \sum_{h0}^{H-1} grad_k_{t,i,h} \cdot V_{t,i,h} $$第 3 步V的总梯度及最终输出梯度$$ grad_V_{t,i,h} grad_output_{t,h} \cdot probs_{t,i} grad_k_{t,i,h} \cdot inv_norm_{t,i} - \frac{grad_inv_norm_{t,i} \cdot inv_norm_{t,i}^{3}}{H} \cdot V_{t,i,h} $$$$ grad_block_res_{t,i,h} grad_V_{t,i,h}, \quad i N $$$$ grad_partial_block_{t,h} grad_V_{t,N,h} $$$$ grad_norm_weight_{h} grad_score_weight_{h} \cdot proj_weight_{0,h} $$$$ grad_proj_weight_{0,h} grad_score_weight_{h} \cdot norm_weight_{h} $$从公式结构可以清晰看到融合反向的三大组成部分注意力加权求和反向grad_V的第一项grad_output × probs对应out Σ probs·V对V的直接梯度RMS 归一化反向grad_V的第三项携带 $inv_norm^3/H$ 修正因子是 RMSNorm 反向传播的解析梯度形式权重梯度grad_norm_weight、grad_proj_weight由grad_score_weight分别乘以对方权重得到对应前向score_weight norm_weight ⊙ proj_weight的逐元素乘积求导。需要特别强调的是当前版本直接使用前向保存的probs和invNorm进行反向计算不根据validBlockNum重新构造掩码因此反向结果完全以保存的中间量为准。参数说明输入、属性与输出总览以下参数表来自算子 README.md 的参数说明章节参数名输入/输出/属性描述数据类型数据格式partialBlock输入前向输入前缀和拼接后作为第 $N1$ 个 valueshape 为 $[T,H]$FLOAT16、BFLOAT16、FLOAT32NDblockRes输入前向输入分块残差拼接后作为前 $N$ 个 valueshape 为 $[T,N,H]$FLOAT16、BFLOAT16、FLOAT32NDprojWeight输入前向投影权重与 normWeight 共同构成 score_weightshape 为 $[1,H]$FLOAT16、BFLOAT16、FLOAT32NDnormWeight输入前向归一化权重与 projWeight 共同构成 score_weightshape 为 $[H]$FLOAT16、BFLOAT16、FLOAT32NDgradHiddenStates输入前向输出 out 的上游梯度shape 为 $[T,H]$FLOAT16、BFLOAT16、FLOAT32NDinvNorm输入前向保存的逐行归一化系数shape 为 $[T,N1]$FLOAT32NDprobs输入前向 softmax 输出概率shape 为 $[T,N1]$FLOAT32NDvalidBlockNum属性预留属性默认值为 -1当前不参与计算仅支持传入 -1 或 NINT64-gradPartialBlock输出partialBlock 的梯度shape 与 partialBlock 一致同主输入NDgradBlockRes输出blockRes 的梯度shape 与 blockRes 一致同主输入NDgradProjWeight输出projWeight 的梯度shape 与 projWeight 一致同主输入NDgradNormWeight输出normWeight 的梯度shape 与 normWeight 一致同主输入ND在 Host 侧算子定义 block_attention_residuals_grad_def.cpp 中7 个输入均注册为REQUIRED且带有AutoContiguous()标记这与接口输入支持非连续 Tensor内部自动转 Contiguous的行为一致valid_block_num注册为OPTIONAL的 INT64 属性且默认值为 -1。aclnn 接口参数细节在 aclnn 接口文档 aclnnBlockAttentionResidualsGrad.md 中每个输入张量还给出了更精确的使用说明partialBlock数据类型与其余输入保持一致shape(T,H)支持非连续 TensorblockRes数据类型与 partialBlock 保持一致shape(T,N,H)支持非连续 TensorprojWeight数据类型与 partialBlock 保持一致shape(1,H)支持非连续 TensornormWeight数据类型与 partialBlock 保持一致shape(H)支持非连续 TensorgradHiddenStates数据类型与 partialBlock 保持一致shape(T,H)支持非连续 TensorinvNorm仅支持 FLOAT32shape(T,N1)支持非连续 Tensorprobs仅支持 FLOAT32shape(T,N1)支持非连续 TensorvalidBlockNum预留属性当前不参与计算仅支持传入 -1四个输出张量数据类型与 shape 分别与对应主输入保持一致gradPartialBlock↔partialBlock、gradBlockRes↔blockRes、gradProjWeight↔projWeight、gradNormWeight↔normWeightworkspaceSize返回需要在 Device 侧申请的 workspace 大小executor返回包含算子计算流程的 op 执行器。约束说明根据 README.md 与 aclnn 接口文档算子约束如下$T \ge 1$$0 \le N \le 128$$H \ge 1$各张量中的 $T$、$H$ 以及invNorm/probs的第 2 维 $N1$ 需保持一致主输入partialBlock/blockRes/projWeight/normWeight/gradHiddenStates支持 FLOAT16、BFLOAT16、FLOAT32且dtype 需一致invNorm/probs仅支持 FLOAT32输入支持非连续 Tensor接口内部会先转为 Contiguous 再计算输出 dtype 与对应主输入保持一致validBlockNum为预留属性不同取值不影响当前版本的计算结果仅支持传入 -1 或 NaclnnBlockAttentionResidualsGrad默认确定性实现确定性计算对数值复现和调试非常友好。边界场景的源码级说明aclnn 第一段接口实现 aclnn_block_attention_residuals_grad.cpp 中对约束做了完整落地并在HandleEmptyTensorL292-L329中处理了边界张量场景T0 或 H0 时aclnn 层跳过主算子非空的权重梯度输出清零空输出保持对应输入 shape。H0 时所有输出均为空不安排清零任务T0 且 H0 时的清零任务仍需调用第二阶段接口执行仅 N0 且 T、H 非零时不提前返回仍计算 partialBlock 及权重梯度此时只有前缀和这一个 value即退化为无 block 残差的纯加权路径。在 Host Tiling 的 shape 校验函数CheckShapeBlockAttentionResidualsGrad见 block_attention_residuals_grad_tiling.cpp L271-L341中还会强制校验partialBlock必须是 2 维、blockRes必须是 3 维、B与blockRes.shape[0]一致、H与blockRes.shape[2]一致、N在[0, 128]内超出 128 报错因为 K 轴 meta Buffer 按totalBlocks N1驻留 UB超过设计上限。调用方式aclnn 两段式接口aclnnBlockAttentionResidualsGrad采用 CANN aclnn 标准的两段式接口必须先调用GetWorkspaceSize阶段接口获取计算所需 workspace 大小以及包含算子计算流程的执行器再调用执行接口完成计算。函数原型aclnnStatus aclnnBlockAttentionResidualsGradGetWorkspaceSize( const aclTensor *partialBlock, const aclTensor *blockRes, const aclTensor *projWeight, const aclTensor *normWeight, const aclTensor *gradHiddenStates, const aclTensor *invNorm, const aclTensor *probs, int64_t validBlockNum, const aclTensor *gradPartialBlock, const aclTensor *gradBlockRes, const aclTensor *gradProjWeight, const aclTensor *gradNormWeight, uint64_t *workspaceSize, aclOpExecutor **executor);aclnnStatus aclnnBlockAttentionResidualsGrad( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);返回值与错误码aclnnBlockAttentionResidualsGradGetWorkspaceSize第一段接口完成入参校验出现以下场景时报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001partialBlock、blockRes、projWeight、normWeight、gradHiddenStates、invNorm、probs 及输出张量存在空指针ACLNN_ERR_PARAM_INVALID161002输入张量的数据类型、数据格式或 shape 不在支持的范围内ACLNN_ERR_RUNTIME_ERROR361001API 内存调用 npu runtime 的接口异常第二段接口aclnnBlockAttentionResidualsGrad的参数为workspaceDevice 侧申请的 workspace 内存地址、workspaceSize由第一段接口获取的大小、executor包含算子计算流程的执行器、stream指定执行任务的 Stream。完整调用示例C 版仓库提供了两个可直接参考的 C 调用样例首个样例 test_aclnn_block_attention_residuals_grad.cpp 完整演示了 aclnn 调用流程核心骨架如下T2、N4、H64 的小规模用例#include iostream #include vector #include cstring #include acl/acl.h #include aclnnop/aclnn_block_attention_residuals_grad.h int main() { int32_t deviceId 0; aclrtStream stream; auto ret Init(deviceId, stream); // aclInit / aclrtSetDevice / aclrtCreateStream const int64_t T 2; const int64_t N 4; const int64_t H 64; const int64_t N1 N 1; std::vectorint64_t partialBlockShape {T, H}; std::vectorint64_t blockResShape {T, N, H}; std::vectorint64_t projWeightShape {1, H}; std::vectorint64_t normWeightShape {H}; std::vectorint64_t gradHiddenStatesShape {T, H}; std::vectorint64_t invNormShape {T, N1}; std::vectorint64_t probsShape {T, N1}; // 主输入用 FP160x3C00 即 1.0invNorm/probs 用 FP32 std::vectoruint16_t partialBlockData(GetShapeSize(partialBlockShape), 0x3C00); std::vectoruint16_t blockResData(GetShapeSize(blockResShape), 0x3C00); std::vectoruint16_t projWeightData(GetShapeSize(projWeightShape), 0x3C00); std::vectoruint16_t normWeightData(GetShapeSize(normWeightShape), 0x3C00); std::vectoruint16_t gradHiddenStatesData(GetShapeSize(gradHiddenStatesShape), 0x3C00); std::vectorfloat invNormData(GetShapeSize(invNormShape), 1.0f); std::vectorfloat probsData(GetShapeSize(probsShape), 1.0f); int64_t validBlockNum -1; // 预留属性仅支持传入-1 // 依次创建输入 aclTensoraclrtMalloc aclrtMemcpy aclCreateTensor // ... 创建 gradPartialBlock / gradBlockRes / gradProjWeight / gradNormWeight 输出张量 ... uint64_t workspaceSize 0; aclOpExecutor* executor; // 第一段获取 workspace 大小与执行器 ret aclnnBlockAttentionResidualsGradGetWorkspaceSize( partialBlock, blockRes, projWeight, normWeight, gradHiddenStates, invNorm, probs, validBlockNum, gradPartialBlock, gradBlockRes, gradProjWeight, gradNormWeight, workspaceSize, executor); // 申请 workspace若 workspaceSize 0 void* workspaceAddr nullptr; if (workspaceSize 0) { aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 第二段执行计算 ret aclnnBlockAttentionResidualsGrad(workspaceAddr, workspaceSize, executor, stream); // 同步等待 aclrtSynchronizeStream(stream); // ... 资源释放aclDestroyTensor / aclrtFree / aclrtDestroyStream / aclrtResetDevice / aclFinalize return 0; }大 H 场景SPLIT_H 模板示例第二个样例 test_aclnn_block_attention_residuals_grad_split_h.cpp 专门演示大 H 场景下 SPLIT_H 模板的调用方式。示例注释说明在 Ascend 910B 上使用 FP16、K3 时若 H8192则超过 FULL_H 路径的 UB 容量Kernel 会切换到 SPLIT_H 路径日志输出hMode1, kernelSPLIT_H。该样例还演示了更有区分度的数据构造——让三个 V block 取不同值FP16 的 0.5、1.5 与 1.0并配合非均匀 probs0.2/0.3/0.5以产生非零的gradScore和varianceScale便于验证反向数值。两个样例的调用方式完全相同两段式接口区别仅在于 shape 规模触发了不同的 Kernel 分派路径这说明SPLIT_H 对上层调用完全透明。PyTorch 侧封装block_attention_residuals_backward除 aclnn C 接口外仓库还提供了 PyTorch 扩展封装见 torch_extension/block_attention_residuals_backward.py。其对外函数签名与 aclnn 接口一一对应def block_attention_residuals_backward( partial_block: torch.Tensor, # [T, H]FP16/BF16/FP32 block_res: torch.Tensor, # [T, N, H] proj_weight: torch.Tensor, # [1, H] norm_weight: torch.Tensor, # [H] grad_hidden_states: torch.Tensor, # [T, H] inv_norm: torch.Tensor, # [T, N1]仅 FP32 probs: torch.Tensor, # [T, N1]仅 FP32 *, valid_block_num: int -1, ) - Tuple[torch.Tensor, ...]:该封装通过torch.library.impl注册了名为block_attention_residuals_backward的自定义算子schema 见 block_attention_residuals_backward.py L253-L259返回(Tensor, Tensor, Tensor, Tensor)并依次完成四类入参校验_check_dimensions各张量必须是规定维度partialBlock 2 维、blockRes 3 维、projWeight 2 维、normWeight 1 维、其余 2 维_check_dtypes主输入须为 FP16/BF16/FP32 且与partial_block一致inv_norm/probs必须为 FP32_check_shapestoken 数 ≥ 0、block_num在[0, 128]、各张量间 T/H/N1 关系一致、valid_block_num为 -1 或block_res.size(1)_check_device所有输入必须在同一 Device 上。PyTorch 侧 docstring 明确指出该算子融合了 softmax 反向、RMS 归一化反向与注意力加权求和反向产生 BlockAttentionResiduals 前向算子保存的四个梯度。源码级实现原理Host 侧Tiling 与 Workspace 规划Tiling 实现位于 block_attention_residuals_grad_tiling.cpp其核心逻辑可概括为平台信息获取TilingPrepareForBlockAttentionResidualsGradL506-L521通过platform_ascendc获取 AIV 核数coreNum与 UB 大小ubSize写入编译期信息结构shape/dtype 校验L271-L405校验维度、B/H/N 一致性、N ≤ 128、主输入 dtype 一致、invNorm/probs 为 FP32H 轴切分决策CalcHiddenTilingL226-L249根据架构regbase 平台走arch35否则走arch22分别估算 FULL_H 路径所需 UB 字节数若requiredFull availableUb则判定为 SPLIT_H并依据线性 UB 模型计算出最大可容纳的 H tile 大小hiddenTileSize最终设置对应的 TilingKeyTPL_H_MODE_SPLIT或TPL_H_MODE_FULLWorkspace 分配CalcWorkspaceSizeL446-L467每个核预留AlignUp(H × sizeof(float), 512B)的 H 轴归约空间SPLIT_H 模式下额外增加两份[B, N1]的 FP32 元数据空间分别保存gradScore与varianceScale用于跨 H tile 的二次归约。从 Tiling 代码注释可以看到arch22 与 arch35 在 SPLIT_H 路径的 Buffer 构成上存在差异arch22 多两个 K 轴 Kahan 补偿 Bufferarch35 多一个写 Workspace 的 FP32 Buffer 等这解释了为什么 tiling 逻辑按架构分别建模。Kernel 侧模板分派Kernel 入口 block_attention_residuals_grad.cpp 使用if constexpr (hMode ...)做编译期模板分派TPL_H_MODE_FULL调用BlockAttentionResidualsGradDTYPE_PARTIAL_BLOCK见 arch22/block_attention_residuals_grad_kernel.hTPL_H_MODE_SPLIT调用BlockAttentionResidualsGradSplitHDTYPE_PARTIAL_BLOCK见 arch22/block_attention_residuals_grad_split_h_arch22.h。Kernel 通过REGISTER_TILING_DEFAULT与GET_TILING_DATA_WITH_STRUCT读取 Host Tiling 下发的BlockAttentionResidualsGradTilingDatabatchSize、numBlocks、totalBlocks、hiddenSize、hiddenTileSize、hiddenTileNum、coreNum、perCoreWkspBytes、gradScoresWkspOff、varianceScaleWkspOff 等字段从而确定每个核负责的 batch 区间与 H tile 划分。ascend950 平台则走 regbase 实现 arch35/block_attention_residuals_grad_regbase.h即上文 op_def 中opFile.value block_attention_residuals_grad_apt对应的入口。编译与运行在 ops-transformer 仓库根目录执行以下命令即可编译该算子并运行示例示例以ascend950平台与 custom vendor 为例# 在 ops-transformer 仓库根目录执行 bash build.sh --pkg --socascend950 --opsblock_attention_residuals_grad bash build.sh --run_example block_attention_residuals_grad eager cust --socascend950 --vendor_namecustom第一条命令按指定 SoC 编译打包该算子第二条命令编译并运行该算子的 eager 示例。实际运行时请根据目标设备将--soc替换为对应的 SoC 型号如 Atlas A2/A3 系列对应的 SoC。小结BlockAttentionResidualsGrad是 CANN ops-transformer 中注意力残差融合模块的反向基石通过复用前向保存的invNorm与probs在一个 Kernel 内完成 Softmax 反向、RMS 归一化反向与加权求和反向的融合计算并支持 FULL_H/SPLIT_H 两条 H 轴切分路径以覆盖大 hidden size 场景。本文给出的公式推导、参数约束、aclnn 两段式调用示例、PyTorch 封装以及 Host/Kernel 实现要点均可在仓库mhc/block_attention_residuals_grad/目录下找到对应源码佐证可作为二次开发与性能分析的直接参考。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表