CANN ops-transformer 中的 BlockAttentionResidualsGrad:融合 Softmax 与 RMSNorm 反向的注意力残差梯度算子解析
2026/9/20 22:26:36 网站建设 项目流程

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(注意力残差)的反向传播算子。正向前向算子将partialBlockblockRes按 block 维拼接成 value 序列后,依次完成 RMS 归一化、normWeightprojWeight逐元素乘积投影打分、Softmax 归一化以及加权求和融合,输出hiddenStates;当needBackward为 true 时,还会输出反向所需的两份中间量:

  • invNorm:前向逐行 RMS 归一化系数,shape 为[T, N+1],仅 FLOAT32;
  • probs:前向 Softmax 输出概率,shape 为[T, N+1],仅 FLOAT32。

反向算子BlockAttentionResidualsGrad正是利用这两份前向保存的中间量,结合上游梯度gradHiddenStates与权重projWeightnormWeight,一次性完成 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,$N+1$ 个 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 步:outVprobs的反向梯度

$$ g_{t,i} = \sum_{h=0}^{H-1} grad_output_{t,h} \cdot V_{t,i,h} $$

$$ grad_score_{t,i} = probs_{t,i} \cdot \left(g_{t,i} - \sum_{j=0}^{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_{t=0}^{T-1}\sum_{i=0}^{N} grad_score_{t,i} \cdot k_{t,i,h} $$

$$ grad_inv_norm_{t,i} = \sum_{h=0}^{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} $$

从公式结构可以清晰看到融合反向的三大组成部分:

  1. 注意力加权求和反向grad_V的第一项grad_output × probs对应out = Σ probs·VV的直接梯度;
  2. RMS 归一化反向grad_V的第三项携带 $inv_norm^3/H$ 修正因子,是 RMSNorm 反向传播的解析梯度形式;
  3. 权重梯度grad_norm_weightgrad_proj_weightgrad_score_weight分别乘以对方权重得到,对应前向score_weight = norm_weight ⊙ proj_weight的逐元素乘积求导。

需要特别强调的是:当前版本直接使用前向保存的probsinvNorm进行反向计算,不根据validBlockNum重新构造掩码,因此反向结果完全以保存的中间量为准。

参数说明

输入、属性与输出总览

以下参数表来自算子 README.md 的参数说明章节:

参数名输入/输出/属性描述数据类型数据格式
partialBlock输入前向输入前缀和,拼接后作为第 $N+1$ 个 value,shape 为 $[T,H]$FLOAT16、BFLOAT16、FLOAT32ND
blockRes输入前向输入分块残差,拼接后作为前 $N$ 个 value,shape 为 $[T,N,H]$FLOAT16、BFLOAT16、FLOAT32ND
projWeight输入前向投影权重,与 normWeight 共同构成 score_weight,shape 为 $[1,H]$FLOAT16、BFLOAT16、FLOAT32ND
normWeight输入前向归一化权重,与 projWeight 共同构成 score_weight,shape 为 $[H]$FLOAT16、BFLOAT16、FLOAT32ND
gradHiddenStates输入前向输出 out 的上游梯度,shape 为 $[T,H]$FLOAT16、BFLOAT16、FLOAT32ND
invNorm输入前向保存的逐行归一化系数,shape 为 $[T,N+1]$FLOAT32ND
probs输入前向 softmax 输出概率,shape 为 $[T,N+1]$FLOAT32ND
validBlockNum属性预留属性,默认值为 -1,当前不参与计算;仅支持传入 -1 或 NINT64-
gradPartialBlock输出partialBlock 的梯度,shape 与 partialBlock 一致同主输入ND
gradBlockRes输出blockRes 的梯度,shape 与 blockRes 一致同主输入ND
gradProjWeight输出projWeight 的梯度,shape 与 projWeight 一致同主输入ND
gradNormWeight输出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),支持非连续 Tensor;
  • blockRes:数据类型与 partialBlock 保持一致,shape(T,N,H),支持非连续 Tensor;
  • projWeight:数据类型与 partialBlock 保持一致,shape(1,H),支持非连续 Tensor;
  • normWeight:数据类型与 partialBlock 保持一致,shape(H),支持非连续 Tensor;
  • gradHiddenStates:数据类型与 partialBlock 保持一致,shape(T,H),支持非连续 Tensor;
  • invNorm:仅支持 FLOAT32,shape(T,N+1),支持非连续 Tensor;
  • probs:仅支持 FLOAT32,shape(T,N+1),支持非连续 Tensor;
  • validBlockNum:预留属性,当前不参与计算,仅支持传入 -1;
  • 四个输出张量:数据类型与 shape 分别与对应主输入保持一致(gradPartialBlock↔partialBlock、gradBlockRes↔blockRes、gradProjWeight↔projWeight、gradNormWeight↔normWeight);
  • workspaceSize:返回需要在 Device 侧申请的 workspace 大小;
  • executor:返回包含算子计算流程的 op 执行器。

约束说明

根据 README.md 与 aclnn 接口文档,算子约束如下:

  • $T \ge 1$,$0 \le N \le 128$,$H \ge 1$;各张量中的 $T$、$H$ 以及invNorm/probs的第 2 维 $N+1$ 需保持一致;
  • 主输入partialBlock/blockRes/projWeight/normWeight/gradHiddenStates支持 FLOAT16、BFLOAT16、FLOAT32,且dtype 需一致invNorm/probs仅支持 FLOAT32;
  • 输入支持非连续 Tensor,接口内部会先转为 Contiguous 再计算;
  • 输出 dtype 与对应主输入保持一致;
  • validBlockNum为预留属性,不同取值不影响当前版本的计算结果,仅支持传入 -1 或 N;
  • aclnnBlockAttentionResidualsGrad默认确定性实现(确定性计算对数值复现和调试非常友好)。

边界场景的源码级说明

aclnn 第一段接口实现 aclnn_block_attention_residuals_grad.cpp 中对约束做了完整落地,并在HandleEmptyTensor(L292-L329)中处理了边界张量场景:

  • T=0 或 H=0 时:aclnn 层跳过主算子;非空的权重梯度输出清零,空输出保持对应输入 shape。H=0 时所有输出均为空,不安排清零任务;T=0 且 H>0 时的清零任务仍需调用第二阶段接口执行;
  • 仅 N=0 且 T、H 非零时:不提前返回,仍计算 partialBlock 及权重梯度(此时只有前缀和这一个 value,即退化为无 block 残差的纯加权路径)。

在 Host Tiling 的 shape 校验函数CheckShapeBlockAttentionResidualsGrad(见 block_attention_residuals_grad_tiling.cpp L271-L341)中,还会强制校验:partialBlock必须是 2 维、blockRes必须是 3 维、BblockRes.shape[0]一致、HblockRes.shape[2]一致、N[0, 128]内(超出 128 报错,因为 K 轴 meta Buffer 按totalBlocks = N+1驻留 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的参数为:workspace(Device 侧申请的 workspace 内存地址)、workspaceSize(由第一段接口获取的大小)、executor(包含算子计算流程的执行器)、stream(指定执行任务的 Stream)。

完整调用示例:C++ 版

仓库提供了两个可直接参考的 C++ 调用样例,首个样例 test_aclnn_block_attention_residuals_grad.cpp 完整演示了 aclnn 调用流程,核心骨架如下(T=2、N=4、H=64 的小规模用例):

#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::vector<int64_t> partialBlockShape = {T, H}; std::vector<int64_t> blockResShape = {T, N, H}; std::vector<int64_t> projWeightShape = {1, H}; std::vector<int64_t> normWeightShape = {H}; std::vector<int64_t> gradHiddenStatesShape = {T, H}; std::vector<int64_t> invNormShape = {T, N1}; std::vector<int64_t> probsShape = {T, N1}; // 主输入用 FP16(0x3C00 即 1.0),invNorm/probs 用 FP32 std::vector<uint16_t> partialBlockData(GetShapeSize(partialBlockShape), 0x3C00); std::vector<uint16_t> blockResData(GetShapeSize(blockResShape), 0x3C00); std::vector<uint16_t> projWeightData(GetShapeSize(projWeightShape), 0x3C00); std::vector<uint16_t> normWeightData(GetShapeSize(normWeightShape), 0x3C00); std::vector<uint16_t> gradHiddenStatesData(GetShapeSize(gradHiddenStatesShape), 0x3C00); std::vector<float> invNormData(GetShapeSize(invNormShape), 1.0f); std::vector<float> probsData(GetShapeSize(probsShape), 1.0f); int64_t validBlockNum = -1; // 预留属性,仅支持传入-1 // 依次创建输入 aclTensor(aclrtMalloc + 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、K=3 时,若 H=8192,则超过 FULL_H 路径的 UB 容量,Kernel 会切换到 SPLIT_H 路径(日志输出hMode=1, kernel=SPLIT_H)。该样例还演示了更有区分度的数据构造——让三个 V block 取不同值(FP16 的 0.5、1.5 与 1.0),并配合非均匀 probs(0.2/0.3/0.5)以产生非零的gradScorevarianceScale,便于验证反向数值。

两个样例的调用方式完全相同(两段式接口),区别仅在于 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, N+1],仅 FP32 probs: torch.Tensor, # [T, N+1],仅 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_shapes:token 数 ≥ 0、block_num[0, 128]、各张量间 T/H/N+1 关系一致、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,其核心逻辑可概括为:

  1. 平台信息获取TilingPrepareForBlockAttentionResidualsGrad,L506-L521):通过platform_ascendc获取 AIV 核数coreNum与 UB 大小ubSize,写入编译期信息结构;
  2. shape/dtype 校验(L271-L405):校验维度、B/H/N 一致性、N ≤ 128、主输入 dtype 一致、invNorm/probs 为 FP32;
  3. H 轴切分决策CalcHiddenTiling,L226-L249):根据架构(regbase 平台走arch35,否则走arch22)分别估算 FULL_H 路径所需 UB 字节数;若requiredFull > availableUb则判定为 SPLIT_H,并依据线性 UB 模型计算出最大可容纳的 H tile 大小hiddenTileSize,最终设置对应的 TilingKey(TPL_H_MODE_SPLITTPL_H_MODE_FULL);
  4. Workspace 分配CalcWorkspaceSize,L446-L467):每个核预留AlignUp(H × sizeof(float), 512B)的 H 轴归约空间;SPLIT_H 模式下额外增加两份[B, N+1]的 FP32 元数据空间,分别保存gradScorevarianceScale,用于跨 H tile 的二次归约。

从 Tiling 代码注释可以看到,arch22 与 arch35 在 SPLIT_H 路径的 Buffer 构成上存在差异(arch22 多两个 K 轴 Kahan 补偿 Buffer,arch35 多一个写 Workspace 的 FP32 Buffer 等),这解释了为什么 tiling 逻辑按架构分别建模。

Kernel 侧:模板分派

Kernel 入口 block_attention_residuals_grad.cpp 使用if constexpr (hMode == ...)做编译期模板分派:

  • TPL_H_MODE_FULL:调用BlockAttentionResidualsGrad<DTYPE_PARTIAL_BLOCK>(见 arch22/block_attention_residuals_grad_kernel.h);
  • TPL_H_MODE_SPLIT:调用BlockAttentionResidualsGradSplitH<DTYPE_PARTIAL_BLOCK>(见 arch22/block_attention_residuals_grad_split_h_arch22.h)。

Kernel 通过REGISTER_TILING_DEFAULTGET_TILING_DATA_WITH_STRUCT读取 Host Tiling 下发的BlockAttentionResidualsGradTilingData(batchSize、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 --soc=ascend950 --ops=block_attention_residuals_grad bash build.sh --run_example block_attention_residuals_grad eager cust --soc=ascend950 --vendor_name=custom

第一条命令按指定 SoC 编译打包该算子;第二条命令编译并运行该算子的 eager 示例。实际运行时请根据目标设备将--soc替换为对应的 SoC 型号(如 Atlas A2/A3 系列对应的 SoC)。

小结

BlockAttentionResidualsGrad是 CANN ops-transformer 中"注意力残差"融合模块的反向基石,通过复用前向保存的invNormprobs,在一个 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询