
CANN ops-transformer 算子解析MhcPreBackward 反向梯度算子实现与 aclnn 调用指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本指南以 mhc/mhc_pre_backward/README.md 为核心结合 CANN ops-transformer 开源仓库中该算子的 host 侧定义、infershape、tiling 与 op_api 源码系统讲解 MhcPreBackward 算子的功能、反向传播数学原理、参数规格、两段式 aclnn 接口调用方法以及跨芯片架构的实现差异。读完本文你将能够理解 mHCManifold-Constrained Hyper-Connections结构中反向梯度算子的完整数据流并掌握通过aclnnMhcPreBackward与aclnnMhcPreBackwardV2在 Ascend NPU 上实际调用该算子的能力。一、算子定位mHC 超连接结构中的反向梯度计算MhcPreBackward是MhcPre的反向算子二者共同构成 mHCManifold-Constrained Hyper-Connections超连接结构在 Transformer 模型中的前向-反向计算闭环。前向算子MhcPre参见 mhc/mhc_pre/README.md基于一系列计算得到 mHC 架构中的 $H^{res}$ 和 $H^{post}$ 投影矩阵以及 Attention 或 MLP 层的输入矩阵 $h^{in}$。反向算子MhcPreBackward则根据前向缓存与上层回传的梯度计算出对x、phi、alpha、bias等参数的梯度用于网络的反向传播更新。算子主要输出为gradX、gradPhi、gradAlpha、gradBias并且在gamma ! nullptr时额外输出gradGamma。其计算依赖前向阶段缓存下来的invRms、hMix、hPre、hPost四个中间结果同时支持两个可选输入RMSNorm 缩放因子gamma与来自后续路径的gradXPostOptional用于融合 mhc_post 反向输出的 grad_x 累加项。反向计算公式总览$$ \begin{aligned} gradX \nabla_{x}(\text{MhcPre}(x, \phi, \alpha, \gamma)) \ gradPhi \nabla_{\phi}(\text{MhcPre}(x, \phi, \alpha, \gamma)) \ gradAlpha \nabla_{\alpha}(\text{MhcPre}(x, \phi, \alpha, \gamma)) \ gradBias \nabla_{bias}(\text{MhcPre}(x, \phi, \alpha, \gamma)) \end{aligned} $$二、产品支持情况产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×从算子注册源码 mhc_pre_backward_def.cpp 可以印证这一支持矩阵该算子通过OpAICoreConfig分别注册了ascend950对应 Ascend 950PR/950DT 与 Atlas A3 系列以及ascend910b、ascend910_93对应 Atlas A2 训练/推理系列三种 AICore 配置未注册其它平台。三、反向传播的数学原理结合 aclnnMhcPreBackward.md 中给出的正向-反向成对公式可以完整还原 MhcPreBackward 的梯度计算链条。整个反向过程按前向计算的逆序分为 8 个阶段。3.1 输出组合梯度计算hIn 分支前向中 $H_{in}$ 由残差维度加权求和得到$$ H_in \sum_{i1}^{N} x[{B,S,i,:}] \cdot H_pre_n[B,S,i] $$反向计算 $h_{pre}$ 的梯度与 $x$ 的第三条梯度分量$$ \begin{aligned} H_pre_grad \text{Reduce}\left(H_in_grad.\text{unsqueeze}(-2) \odot x, \text{dim}-1\right) \quad ([B,S,N]) \ x_grad_vec3 H_in_grad \times H_pre \quad ([B,S,N,D]) \end{aligned} $$3.2 Sigmoid 门控反向H_pre前向$H_pre \text{Sigmoid}(\alpha_pre \cdot H_pre_1 bias_pre) hc_eps$反向利用 sigmoid 导数 $s(1-s)$$$ \begin{aligned} s H_pre - hc_eps \ H_pre_2_grad H_pre_grad \odot s \odot (1 - s) \ H_pre_1_grad H_pre_2_grad \cdot \alpha_pre \ \alpha_pre_grad \sum_{b,s,n}^{B,S,N} \left(H_pre_2_grad \cdot H_pre_1\right) \ bias_pre_grad \sum_{b,s}^{B,S} H_pre_2_grad \quad ([N]) \end{aligned} $$3.3 Sigmoid 门控反向H_post前向$H_post \text{Sigmoid}(\alpha_post \cdot H_post_1 bias_post) \cdot 2$注意此处系数 2 引入了 $1 - H_post/2$ 形式的导数项$$ \begin{aligned} H_post_2_grad H_post_grad \odot \left(H_post \cdot \left(1 - \frac{H_post}{2}\right)\right) \ H_post_1_grad H_post_2_grad \cdot \alpha_{post} \ \alpha_{post_grad} \sum_{b,s,n}^{B,S,N} \left(H_post_2_grad \cdot H_{post_1}\right) \ bias_post_grad \sum_{b,s}^{B,S} H_post_2_grad \quad ([N]) \end{aligned} $$3.4 残差连接反向H_res前向$H_res \alpha_res \cdot H_res_1 bias_res$反向$$ \begin{aligned} H_res_2_grad H_res_grad \cdot \alpha_{res} \quad ([B,S,N,N]) \ \alpha_res_grad \sum_{b,s,i,j}^{B,S,N,N} \left(H_res_grad \cdot H_res_2\right) \ bias_res_grad \sum_{b,s}^{B,S} H_res_grad \quad ([N,N]) \ H_res_1_grad \text{Reshape}(H_res_2_grad) \quad ([B,S,N^2]) \end{aligned} $$3.5 RMSNorm Fusion 反向前向$H_mix_tmp H_mix \cdot inv_rms$反向先将三段门控梯度拼接再分别乘inv_rms并累加得到inv_rms梯度$$ \begin{aligned} H_mix_tmp_grad \text{Concat}(H_pre_1_grad, H_post_1_grad, H_res_1_grad) \quad ([B,S,2NN^2]) \ H_mix_grad H_mix_tmp_grad \cdot inv_rms \ inv_rms_{grad} \sum_{\text{last_dim}} \left(H_mix_tmp_grad \cdot H_mix\right) \quad ([B,S,1]) \end{aligned} $$3.6 矩阵乘法反向x phiᵀ前向$H_mix x_rs phi^T$其中 $x_rs x \cdot gamma$反向分别对矩阵乘的两个因子求梯度$$ \begin{aligned} x_rs_grad H_mix_grad phi \quad ([B,S,ND]) \ X \text{Reshape}(x_rs, [B\cdot S, ND]) \ G \text{Reshape}(H_mix_grad, [B\cdot S, 2NN^2]) \ phi_{grad} G^T X \quad ([2NN^2, ND]) \end{aligned} $$3.7 特征缩放反向gamma与 RMS 归一化梯度特征缩放反向$$ \begin{aligned} x_grad_mm x_rs_grad \cdot gamma \ gamma_grad \sum_{b1}^{B}\sum_{s1}^{S} (x \cdot x_rs_grad) \quad ([N,D]) \end{aligned} $$前向中 $inv_rms \dfrac{1}{\sqrt{\frac{1}{n}\sum_{i1}^{n}x_i^2 eps}}$其中 $n N \cdot D$其反向为$$ \begin{aligned} x_rs_grad_inv - \left(\frac{inv_rms_grad \cdot {inv_rms}^3}{N\cdot D}\right) \cdot x_rs \ x_rs_grad x_grad_mm x_rs_grad_inv \ x_grad_vec1 \text{Reshape}(x_rs_grad, [B,S,N,D]) \ x_grad x_grad_vec3 x_grad_vec1 \end{aligned} $$3.8 融合 mhc_post 的 grad_x 相加若传入可选的grad_x_post来自 mhc_post 反向路径的 gradX则在最终输出前执行一次累加$$ x_grad x_grad grad_x_post $$四、参数说明下表继承自 README.md 并补充了 aclnn 接口文档 中的 shape 约束便于直接对照构造 Tensor。参数名输入/输出/属性描述数据类型数据格式x输入mHC 层输入数据shape 为 (B,S,N,D) 或 (T,N,D)BFLOAT16、FLOAT16NDphi输入mHC 参数矩阵shape 为 (2NN·N, N·D) 或 (2NN!, N·D)FLOAT32NDalpha输入mHC 缩放参数shape 为 (3)FLOAT32NDgrad_h_in输入对 h_in 的梯度shape 为 (B,S,D) 或 (T,D)BFLOAT16、FLOAT16NDgrad_h_post输入对 h_post 的梯度shape 为 (B,S,N) 或 (T,N)FLOAT32NDgrad_h_res输入对 h_res 的梯度shape 为 (B,S,N,N)/(B,S,N!)/(T,N,N)/(T,N!)FLOAT32NDinv_rms输入前向缓存的 inv_rmsshape 为 (B,S) 或 (T)FLOAT32NDh_mix输入前向缓存的 h_mixshape 为 (B,S,2NN·N)/(B,S,2NN!)/(T,2NN·N)/(T,2NN!)FLOAT32NDh_pre输入前向缓存的 h_preshape 为 (B,S,N) 或 (T,N)FLOAT32NDh_post输入前向缓存的 h_postshape 为 (B,S,N) 或 (T,N)FLOAT32NDgamma可选输入RMSNorm 缩放因子shape 为 (N,D)传 nullptr 表示全 1FLOAT32NDgrad_x_post可选输入来自后续路径的 grad_x 累加项shape 为 (B,S,N,D) 或 (T,N,D)传 nullptr 表示全 0BFLOAT16、FLOAT16NDhc_eps属性h_pre sigmoid 后使用的 eps 参数建议值 1e-6默认 1e-6FLOAT32-grad_x输出x 的梯度与输入 x 维度、类型一致BFLOAT16、FLOAT16NDgrad_phi输出phi 的梯度与输入 phi 的 shape 一致FLOAT32NDgrad_alpha输出alpha 的梯度shape 为 (3)FLOAT32NDgrad_bias输出bias 整体梯度shape 为 (2NN·N) 或 (2NN!)FLOAT32NDgrad_gamma可选输出gamma 的梯度shape 为 (N,D)仅当输入 gamma 非 nullptr 时输出FLOAT32ND关于融合维度fusionSize的说明phi 的第 0 维即融合维度。在 infershape 源码 中fusionSize直接取自 phi 第 0 维且注释明确指出其平台相关性Atlas A2ascend910b上为N! 2NA3/A5ascend950上为N² 2N。其中N!表示 N 的阶乘排列数如 N4 时 N!24用于支持带 sinkhorn 排列约束的残差变体。五、约束说明不同平台的规格约束差异较大务必按下表核对后再配置 shapeAscend 950PR/Ascend 950DTN 目前支持 4、6、8。D 支持 1~16384需满足 64 元素对齐。phi、gradHRes、hMix、gradPhi、gradBias 的 shape 仅支持N·N形式的融合维度即 2NN·N。确定性计算默认采用确定性实现。Atlas A2 训练系列产品/Atlas A2 推理系列产品N 目前仅支持 4。D 支持 1~100000需满足 128 元素对齐。fusionSize 支持N!2N与N·N2N两种gradHRes 支持 (B,S,N,N)/(B,S,N!)/(T,N,N)/(T,N!) 四种 shape 组合。上述约束与 tiling 层按芯片架构分发实现的事实相吻合在 mhc_pre_backward_tiling.cpp 中tiling 入口会根据 SoC 版本分流到arch22Ascend 910B/A2与arch35Ascend 950/A3两套独立实现。六、调用说明算子提供两个 aclnn 接口均遵循 两段式接口规范先调用GetWorkspaceSize接口完成入参校验、推导 shape、计算 workspace 大小并生成执行器再调用同名执行接口在指定 stream 上发起计算。调用方式调用样例说明aclnn 调用test_aclnn_mhc_pre_backward.cpp通过 aclnnMhcPreBackward 接口方式调用 MhcPreBackward 算子。aclnn 调用test_aclnn_mhc_pre_backward_v2.cpp通过 aclnnMhcPreBackwardV2 接口指定 Cube 计算模式并调用 MhcPreBackward 算子。6.1 aclnnMhcPreBackward 函数原型aclnnStatus aclnnMhcPreBackwardGetWorkspaceSize( const aclTensor *x, const aclTensor *phi, const aclTensor *alpha, const aclTensor *gradHIn, const aclTensor *gradHPost, const aclTensor *gradHRes, const aclTensor *invRms, const aclTensor *hMix, const aclTensor *hPre, const aclTensor *hPost, const aclTensor *gammaOptional, const aclTensor *gradXPostOptional, float hcEps, const aclTensor *gradX, const aclTensor *gradPhi, const aclTensor *gradAlpha, const aclTensor *gradBias, const aclTensor *gradGamma, uint64_t *workspaceSize, aclOpExecutor **executor) aclnnStatus aclnnMhcPreBackward( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)6.2 aclnnMhcPreBackwardV2显式指定 Cube 计算模式aclnnMhcPreBackwardV2 在aclnnMhcPreBackward基础上新增opImplMode参数INT64用于指定算子在 Cube 单元上的计算精度模式opImplMode0Cube 使用 FP32 模式计算opImplMode1Cube 使用 HF32 模式计算其它取值不支持第一段接口会返回ACLNN_ERR_PARAM_INVALID。对应的 V2 原型在GetWorkspaceSize的参数列表中于hcEps之后、gradX之前插入int64_t opImplMode其余参数与 V1 完全一致。需要说明的是V2 接口仅注册在ascend950A3/A5平台详见 def 源码 中op_impl_mode属性与平台配置的关系这也与接口文档中Atlas A2/A3 系列不支持 V2的声明对应。相应地在前向算子 MhcPre 的约束中Atlas A3/A2 平台上op_impl_mode仅支持配置为 0。6.3 返回值与错误码第一段接口完成入参校验返回aclnnStatus状态码完整枚举见 aclnn 返回码返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001必选参数或者输出是空指针。ACLNN_ERR_PARAM_INVALID161002输入变量 x、phi、gamma、alpha 的数据类型和数据格式不在支持范围内V2 接口额外包括 opImplMode 不为 0 或 1。ACLNN_ERR_RUNTIME_ERROR361001API 内部调用 npu runtime 的接口异常。七、完整调用示例以下示例节选自 test_aclnn_mhc_pre_backward.cpp该文件可直接编译运行完整编译与执行流程请参考 编译与运行样例。示例采用T1024, N4, D512的 shape 组合即 TND 紧凑排布格式。#include iostream #include vector #include numeric #include acl/acl.h #include aclnnop/aclnn_mhc_pre_backward.h #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) // 计算 shape 的元素总数 int64_t GetShapeSize(const std::vectorint64_t shape) { int64_t shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } // 将 device 侧结果拷贝回 host 并打印前 10 个元素区分 BF16 与 FP32 void PrintOutResult(std::vectorint64_t shape, void **deviceAddr, const char *name, size_t elemSize sizeof(float)) { auto size GetShapeSize(shape); size_t copyBytes size * elemSize; if (elemSize 2) { std::vectoruint16_t rawData(size, 0); auto ret aclrtMemcpy(rawData.data(), copyBytes, *deviceAddr, copyBytes, ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(copy result from device to host failed. ERROR: %d\n, ret); return); LOG_PRINT(%s result (first 10 elements):\n, name); for (int64_t i 0; i std::min(size, (int64_t)10); i) { union { uint32_t i; float f; } u; u.i (uint32_t)rawData[i] 16; // BF16 - FP32 显示 LOG_PRINT( [%ld] %f\n, i, u.f); } } else { std::vectorfloat resultData(size, 0); auto ret aclrtMemcpy(resultData.data(), copyBytes, *deviceAddr, copyBytes, ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(copy result from device to host failed. ERROR: %d\n, ret); return); LOG_PRINT(%s result (first 10 elements):\n, name); for (int64_t i 0; i std::min(size, (int64_t)10); i) { LOG_PRINT( [%ld] %f\n, i, resultData[i]); } } } // AscendCL 固定初始化device/context/stream int Init(int32_t deviceId, aclrtContext *context, aclrtStream *stream) { auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed. ERROR: %d\n, ret); return ret); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); ret aclrtCreateContext(context, deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateContext failed. ERROR: %d\n, ret); return ret); ret aclrtSetCurrentContext(*context); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetCurrentContext failed. ERROR: %d\n, ret); return ret); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed. ERROR: %d\n, ret); return ret); return 0; } // 申请 device 内存、拷贝数据并创建 aclTensorND 连续布局 template typename T int CreateAclTensor(const std::vectorT hostData, const std::vectorint64_t shape, void **deviceAddr, aclDataType dataType, aclTensor **tensor) { auto size GetShapeSize(shape) * sizeof(T); auto ret aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed. ERROR: %d\n, ret); return ret); ret aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMemcpy failed. ERROR: %d\n, ret); return ret); std::vectorint64_t strides(shape.size(), 1); for (int64_t i shape.size() - 2; i 0; i--) { strides[i] shape[i 1] * strides[i 1]; } *tensor aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1. device/context/stream 初始化按实际 device 填写 deviceId int32_t deviceId 0; aclrtContext context; aclrtStream stream; auto ret Init(deviceId, context, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 构造输入与输出 shapeTND 紧凑排布N4, D512 std::vectorint64_t xShape {1024, 4, 512}; // T, N, D std::vectorint64_t phiShape {24, 2048}; // 2N N*N, N*D std::vectorint64_t alphaShape {3}; // 固定大小 std::vectorint64_t gradHInShape {1024, 512}; // T, D std::vectorint64_t gradHPostShape {1024, 4}; // T, N std::vectorint64_t gradHResShape {1024, 4, 4}; // T, N, N std::vectorint64_t invRmsShape {1024}; // T std::vectorint64_t hMixShape {1024, 24}; // T, 2N N*N std::vectorint64_t hPreShape {1024, 4}; // T, N std::vectorint64_t hPostShape {1024, 4}; // T, N std::vectorint64_t gammaShape {4, 512}; // N, D std::vectorint64_t gradXShape {1024, 4, 512}; // T, N, D std::vectorint64_t gradPhiShape {24, 2048}; // 2N N*N, N*D std::vectorint64_t gradAlphaShape {3}; std::vectorint64_t gradGammaShape {4, 512}; // N, D std::vectorint64_t gradBiasShape {24}; // 2N N*N std::vectorint64_t gradXPostOptionalShape {1024, 4, 512}; void *xDeviceAddr nullptr, *phiDeviceAddr nullptr, *alphaDeviceAddr nullptr; void *gradHInDeviceAddr nullptr, *gradHPostDeviceAddr nullptr, *gradHResDeviceAddr nullptr; void *invRmsDeviceAddr nullptr, *hMixDeviceAddr nullptr, *hPreDeviceAddr nullptr; void *hPostDeviceAddr nullptr, *gammaDeviceAddr nullptr, *gradXDeviceAddr nullptr; void *gradPhiDeviceAddr nullptr, *gradAlphaDeviceAddr nullptr, *gradBiasDeviceAddr nullptr; void *gradGammaDeviceAddr nullptr, *gradXPostOptionalDeviceAddr nullptr; aclTensor *x nullptr, *phi nullptr, *alpha nullptr, *gradHIn nullptr, *gradHPost nullptr; aclTensor *gradHRes nullptr, *invRms nullptr, *hMix nullptr, *hPre nullptr, *hPost nullptr; aclTensor *gamma nullptr, *gradX nullptr, *gradPhi nullptr, *gradAlpha nullptr; aclTensor *gradBias nullptr, *gradGamma nullptr, *gradXPostOptional nullptr; // host 侧数据输入用 1.0 填充输出用 0 初始化 std::vectorshort xHostData(1024 * 4 * 512, 1.0); std::vectorfloat phiHostData(24 * 2048, 1.0); std::vectorfloat alphaHostData(3, 1.0); std::vectorshort gradHInHostData(1024 * 512, 1.0); std::vectorfloat gradHPostHostData(1024 * 4, 1.0); std::vectorfloat gradHResHostData(1024 * 4 * 4, 1.0); std::vectorfloat invRmsHostData(1024, 1.0); std::vectorfloat hMixHostData(1024 * 24, 1.0); std::vectorfloat hPreHostData(1024 * 4, 1.0); std::vectorfloat hPostHostData(1024 * 4, 1.0); std::vectorfloat gammaHostData(4 * 512, 1.0); std::vectorshort gradXHostData(1024 * 4 * 512, 0); std::vectorfloat gradPhiHostData(24 * 2048, 0); std::vectorfloat gradAlphaHostData(3, 0); std::vectorfloat gradBiasHostData(24, 0); std::vectorfloat gradGammaHostData(4 * 512, 0); std::vectorshort gradXPostOptionalHostData(1024 * 4 * 512, 0); ret CreateAclTensor(xHostData, xShape, xDeviceAddr, aclDataType::ACL_BF16, x); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(phiHostData, phiShape, phiDeviceAddr, aclDataType::ACL_FLOAT, phi); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(alphaHostData, alphaShape, alphaDeviceAddr, aclDataType::ACL_FLOAT, alpha); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradHInHostData, gradHInShape, gradHInDeviceAddr, aclDataType::ACL_BF16, gradHIn); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradHPostHostData, gradHPostShape, gradHPostDeviceAddr, aclDataType::ACL_FLOAT, gradHPost); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradHResHostData, gradHResShape, gradHResDeviceAddr, aclDataType::ACL_FLOAT, gradHRes); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(invRmsHostData, invRmsShape, invRmsDeviceAddr, aclDataType::ACL_FLOAT, invRms); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(hMixHostData, hMixShape, hMixDeviceAddr, aclDataType::ACL_FLOAT, hMix); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(hPreHostData, hPreShape, hPreDeviceAddr, aclDataType::ACL_FLOAT, hPre); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(hPostHostData, hPostShape, hPostDeviceAddr, aclDataType::ACL_FLOAT, hPost); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gammaHostData, gammaShape, gammaDeviceAddr, aclDataType::ACL_FLOAT, gamma); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradXHostData, gradXShape, gradXDeviceAddr, aclDataType::ACL_BF16, gradX); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradPhiHostData, gradPhiShape, gradPhiDeviceAddr, aclDataType::ACL_FLOAT, gradPhi); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradAlphaHostData, gradAlphaShape, gradAlphaDeviceAddr, aclDataType::ACL_FLOAT, gradAlpha); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradBiasHostData, gradBiasShape, gradBiasDeviceAddr, aclDataType::ACL_FLOAT, gradBias); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradGammaHostData, gradGammaShape, gradGammaDeviceAddr, aclDataType::ACL_FLOAT, gradGamma); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gradXPostOptionalHostData, gradXPostOptionalShape, gradXPostOptionalDeviceAddr, aclDataType::ACL_BF16, gradXPostOptional); CHECK_RET(ret ACL_SUCCESS, return ret); float hc_eps 1e-6; // h_pre sigmoid 后的 eps // 3. 两段式调用先获取 workspace 大小与执行器 uint64_t workspaceSize 0; aclOpExecutor *executor; ret aclnnMhcPreBackwardGetWorkspaceSize(x, phi, alpha, gradHIn, gradHPost, gradHRes, invRms, hMix, hPre, hPost, gamma, gradXPostOptional, hc_eps, gradX, gradPhi, gradAlpha, gradBias, gradGamma, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnMhcPreBackwardGetWorkspaceSize failed. ERROR: %d\n, ret); return ret); // 按计算出的 workspaceSize 申请 device 内存 void *workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(allocate workspace failed. ERROR: %d\n, ret); return ret); } // 再执行算子 ret aclnnMhcPreBackward(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnMhcPreBackward failed. ERROR: %d\n, ret); return ret); // 4. 同步等待任务执行结束 ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed. ERROR: %d\n, ret); return ret); // 5. 取回结果并打印BF16 输出按 2 字节处理 PrintOutResult(gradXShape, gradXDeviceAddr, gradX, sizeof(short)); PrintOutResult(gradPhiShape, gradPhiDeviceAddr, gradPhi); PrintOutResult(gradAlphaShape, gradAlphaDeviceAddr, gradAlpha); PrintOutResult(gradBiasShape, gradBiasDeviceAddr, gradBias); PrintOutResult(gradGammaShape, gradGammaDeviceAddr, gradGamma); // 6. 释放 aclTensor 与 device 资源 aclDestroyTensor(x); aclDestroyTensor(phi); aclDestroyTensor(alpha); aclDestroyTensor(gradHIn); aclDestroyTensor(gradHPost); aclDestroyTensor(gradHRes); aclDestroyTensor(invRms); aclDestroyTensor(hMix); aclDestroyTensor(hPre); aclDestroyTensor(hPost); aclDestroyTensor(gamma); aclDestroyTensor(gradX); aclDestroyTensor(gradPhi); aclDestroyTensor(gradAlpha); aclDestroyTensor(gradBias); aclDestroyTensor(gradGamma); aclDestroyTensor(gradXPostOptional); aclrtFree(xDeviceAddr); aclrtFree(phiDeviceAddr); aclrtFree(alphaDeviceAddr); aclrtFree(gradHInDeviceAddr); aclrtFree(gradHPostDeviceAddr); aclrtFree(gradHResDeviceAddr); aclrtFree(invRmsDeviceAddr); aclrtFree(hMixDeviceAddr); aclrtFree(hPreDeviceAddr); aclrtFree(hPostDeviceAddr); aclrtFree(gammaDeviceAddr); aclrtFree(gradXDeviceAddr); aclrtFree(gradPhiDeviceAddr); aclrtFree(gradAlphaDeviceAddr); aclrtFree(gradBiasDeviceAddr); aclrtFree(gradGammaDeviceAddr); aclrtFree(gradXPostOptionalDeviceAddr); if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtDestroyContext(context); aclrtResetDevice(deviceId); aclFinalize(); return 0; }若改用 V2 接口仅需两处调整头文件改为aclnnop/aclnn_mhc_pre_backward_v2.h并在第一段调用中于hc_eps之后传入int64_t opImplMode 1;示例选用 HF32 Cube 模式同时将接口名替换为aclnnMhcPreBackwardV2GetWorkspaceSize/aclnnMhcPreBackwardV2。八、Host 侧实现链路解读8.1 算子定义def在 mhc_pre_backward_def.cpp 中可以看到完整的算子注册信息12 个输入x、phi、alpha、grad_h_in、grad_h_post、grad_h_res、inv_rms、h_mix、h_pre、h_post为必选gamma、grad_x_post为可选OPTIONAL与文档参数表一一对应5 个输出grad_x、grad_phi、grad_alpha、grad_bias必选grad_gamma可选2 个属性hc_epsFLOAT默认1e-6f与op_impl_modeINT默认 0动态能力DynamicCompileStaticFlag(true)、DynamicShapeSupportFlag(true)、DynamicRankSupportFlag(true)支持动态 shape 与动态 rank即 B、S 泛化。8.2 输出 shape 推导infershapeinfershape 实现 的推导逻辑完全围绕梯度 shape 必须与前向一致展开校验维度组合ValidateInputDims要求gradHIn与gradHPost维数相等且gradHRes必须匹配对应格式BSND 时 gradHRes 为 4 维 BSNN 或 3 维 BSN!TND 时 gradHRes 为 3 维 TNN 或 2 维 TN!推导 gradXgradX的 shape 由gradHIn与gradHPost组合得出BS 格式下为 (B, S, N, D)T 格式下为 (T, N, D)推导其余输出gradPhi与gradBias的第 0 维取自phi的第 0 维fusionSizegradAlpha固定为 3gradGamma为 (N, D)推导数据类型gradX继承grad_h_in的类型BF16/FP16其余梯度输出统一为 FLOAT32。8.3 Tiling 分发与平台适配Tiling 入口 mhc_pre_backward_tiling.cpp 根据 SoC 版本将计算任务分发给两套实现arch35Ascend 950/A3/A5通过 TilingRegistry 注册的模板实现kernel 侧包含 Cube 计算见 arch35 目录支持 FP32/HF32 两种 Cube 模式对应 V2 接口的opImplMode参数arch22Ascend 910B/A2独立的TilingMhcPreBackwardArch22实现见 arch22 目录对应 A2 平台仅支持 N4、D 128 对齐的规格。此外TilingPrepare4MhcPreBackward在编译期采集各核 AIC/AIV 数量与 UB/L1/L2/L0 各级缓存大小供 tiling 决策切分策略使用。九、单元测试验证要点算子的 op_api 单测位于 tests/ut/op_api/test_aclnn_mhc_pre_backward.cpp主要覆盖以下行为基础路径N4、D8、T10 的 TND 排布下aclnnMhcPreBackward返回ACLNN_SUCCESS阶乘残差变体RunMhcPreBackward(true)使用fusionSize N! 2NN4 时为 24的 shape验证 A2 平台支持的 N! 形式融合维度V2 模式校验opImplMode1HF32调用成功opImplMode2与-1均返回失败验证了仅支持 0/1的约束空指针校验GetWorkspaceSize传入全空指针时返回ACLNN_ERR_PARAM_NULLPTR161001平台模拟测试类通过op::SetPlatformSocVersion(op::SocVersion::ASCEND950)模拟 A5 平台运行场景并在用例结束后恢复原 SoC 版本。infershape 与 tiling 的 host 侧单测分别位于 tests/ut/op_host/可用于回归验证 shape 推导与 tiling 参数生成逻辑。十、小结MhcPreBackward 是 mHC 超连接结构中负责反向传播的关键算子通过将 RMSNorm、Sigmoid 门控、矩阵乘、残差连接等前向步骤的梯度计算融合进单个 NPU 算子避免了逐层物化中间梯度从而在反向传播阶段获得更好的访存效率。实际使用中需重点核对三点平台对应的 N 值支持范围与 D 对齐约束950 平台 64 元素对齐、A2 平台 128 元素对齐、fusionSize 的平台差异950 平台仅 N²2NA2 平台还支持 N!2N、以及 V2 接口的opImplMode仅在 A3/A5 平台可用。结合本仓库中 MhcPre 前向算子 与同目录下其它 mHC 系列算子如 mhc_post、mhc_pre_sinkhorn 等的文档可以拼出完整的 mHC 结构前向-反向实现全景。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考