news 2026/9/21 18:19:58

CANN ops-transformer 算子解析:MhcPreBackward 反向梯度算子实现与 aclnn 调用指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN ops-transformer 算子解析:MhcPreBackward 反向梯度算子实现与 aclnn 调用指南

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 接口调用方法以及跨芯片架构的实现差异。读完本文,你将能够理解 mHC(Manifold-Constrained Hyper-Connections)结构中反向梯度算子的完整数据流,并掌握通过aclnnMhcPreBackwardaclnnMhcPreBackwardV2在 Ascend NPU 上实际调用该算子的能力。

一、算子定位:mHC 超连接结构中的反向梯度计算

MhcPreBackwardMhcPre的反向算子,二者共同构成 mHC(Manifold-Constrained Hyper-Connections)超连接结构在 Transformer 模型中的前向-反向计算闭环。

前向算子MhcPre(参见 mhc/mhc_pre/README.md)基于一系列计算得到 mHC 架构中的 $H^{res}$ 和 $H^{post}$ 投影矩阵,以及 Attention 或 MLP 层的输入矩阵 $h^{in}$。反向算子MhcPreBackward则根据前向缓存与上层回传的梯度,计算出对xphialphabias等参数的梯度,用于网络的反向传播更新。

算子主要输出为gradXgradPhigradAlphagradBias,并且在gamma != nullptr时额外输出gradGamma。其计算依赖前向阶段缓存下来的invRmshMixhPrehPost四个中间结果,同时支持两个可选输入: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 系列)以及ascend910bascend910_93(对应 Atlas A2 训练/推理系列)三种 AICore 配置,未注册其它平台。

三、反向传播的数学原理

结合 aclnnMhcPreBackward.md 中给出的正向-反向成对公式,可以完整还原 MhcPreBackward 的梯度计算链条。整个反向过程按前向计算的逆序分为 8 个阶段。

3.1 输出组合梯度计算(hIn 分支)

前向中 $H_{in}$ 由残差维度加权求和得到:

$$ H_in = \sum_{i=1}^{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,2N+N^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, 2N+N^2]) \ phi_{grad} &= G^T @ X \quad ([2N+N^2, ND]) \end{aligned} $$

3.7 特征缩放反向(gamma)与 RMS 归一化梯度

特征缩放反向:

$$ \begin{aligned} x_grad_mm &= x_rs_grad \cdot gamma \ gamma_grad &= \sum_{b=1}^{B}\sum_{s=1}^{S} (x \cdot x_rs_grad) \quad ([N,D]) \end{aligned} $$

前向中 $inv_rms = \dfrac{1}{\sqrt{\frac{1}{n}\sum_{i=1}^{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、FLOAT16ND
phi输入mHC 参数矩阵,shape 为 (2N+N·N, N·D) 或 (2N+N!, N·D)FLOAT32ND
alpha输入mHC 缩放参数,shape 为 (3)FLOAT32ND
grad_h_in输入对 h_in 的梯度,shape 为 (B,S,D) 或 (T,D)BFLOAT16、FLOAT16ND
grad_h_post输入对 h_post 的梯度,shape 为 (B,S,N) 或 (T,N)FLOAT32ND
grad_h_res输入对 h_res 的梯度,shape 为 (B,S,N,N)/(B,S,N!)/(T,N,N)/(T,N!)FLOAT32ND
inv_rms输入前向缓存的 inv_rms,shape 为 (B,S) 或 (T)FLOAT32ND
h_mix输入前向缓存的 h_mix,shape 为 (B,S,2N+N·N)/(B,S,2N+N!)/(T,2N+N·N)/(T,2N+N!)FLOAT32ND
h_pre输入前向缓存的 h_pre,shape 为 (B,S,N) 或 (T,N)FLOAT32ND
h_post输入前向缓存的 h_post,shape 为 (B,S,N) 或 (T,N)FLOAT32ND
gamma可选输入RMSNorm 缩放因子,shape 为 (N,D);传 nullptr 表示全 1FLOAT32ND
grad_x_post可选输入来自后续路径的 grad_x 累加项,shape 为 (B,S,N,D) 或 (T,N,D);传 nullptr 表示全 0BFLOAT16、FLOAT16ND
hc_eps属性h_pre sigmoid 后使用的 eps 参数,建议值 1e-6,默认 1e-6FLOAT32-
grad_x输出x 的梯度,与输入 x 维度、类型一致BFLOAT16、FLOAT16ND
grad_phi输出phi 的梯度,与输入 phi 的 shape 一致FLOAT32ND
grad_alpha输出alpha 的梯度,shape 为 (3)FLOAT32ND
grad_bias输出bias 整体梯度,shape 为 (2N+N·N) 或 (2N+N!)FLOAT32ND
grad_gamma可选输出gamma 的梯度,shape 为 (N,D);仅当输入 gamma 非 nullptr 时输出FLOAT32ND

关于融合维度(fusionSize)的说明:phi 的第 0 维即融合维度。在 infershape 源码 中,fusionSize直接取自 phi 第 0 维,且注释明确指出其平台相关性:Atlas A2(ascend910b)上为N! + 2N,A3/A5(ascend950)上为N² + 2N。其中N!表示 N 的阶乘排列数(如 N=4 时 N!=24),用于支持带 sinkhorn 排列约束的残差变体。

五、约束说明

不同平台的规格约束差异较大,务必按下表核对后再配置 shape:

  • Ascend 950PR/Ascend 950DT

    • N 目前支持 4、6、8。
    • D 支持 1~16384,需满足 64 元素对齐。
    • phi、gradHRes、hMix、gradPhi、gradBias 的 shape 仅支持N·N形式的融合维度(即 2N+N·N)。
    • 确定性计算:默认采用确定性实现。
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品

    • N 目前仅支持 4。
    • D 支持 1~100000,需满足 128 元素对齐。
    • fusionSize 支持N!+2NN·N+2N两种;gradHRes 支持 (B,S,N,N)/(B,S,N!)/(T,N,N)/(T,N!) 四种 shape 组合。

上述约束与 tiling 层按芯片架构分发实现的事实相吻合:在 mhc_pre_backward_tiling.cpp 中,tiling 入口会根据 SoC 版本分流到arch22(Ascend 910B/A2)与arch35(Ascend 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 单元上的计算精度模式:

  • opImplMode=0:Cube 使用 FP32 模式计算;
  • opImplMode=1:Cube 使用 HF32 模式计算;
  • 其它取值不支持,第一段接口会返回ACLNN_ERR_PARAM_INVALID

对应的 V2 原型在GetWorkspaceSize的参数列表中于hcEps之后、gradX之前插入int64_t opImplMode,其余参数与 V1 完全一致。需要说明的是,V2 接口仅注册在ascend950(A3/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(该文件可直接编译运行),完整编译与执行流程请参考 编译与运行样例。示例采用T=1024, N=4, D=512的 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::vector<int64_t> &shape) { int64_t shapeSize = 1; for (auto i : shape) { shapeSize *= i; } return shapeSize; } // 将 device 侧结果拷贝回 host 并打印前 10 个元素(区分 BF16 与 FP32) void PrintOutResult(std::vector<int64_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::vector<uint16_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::vector<float> 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 内存、拷贝数据并创建 aclTensor(ND 连续布局) template <typename T> int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_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::vector<int64_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. 构造输入与输出 shape(TND 紧凑排布,N=4, D=512) std::vector<int64_t> xShape = {1024, 4, 512}; // T, N, D std::vector<int64_t> phiShape = {24, 2048}; // 2N + N*N, N*D std::vector<int64_t> alphaShape = {3}; // 固定大小 std::vector<int64_t> gradHInShape = {1024, 512}; // T, D std::vector<int64_t> gradHPostShape = {1024, 4}; // T, N std::vector<int64_t> gradHResShape = {1024, 4, 4}; // T, N, N std::vector<int64_t> invRmsShape = {1024}; // T std::vector<int64_t> hMixShape = {1024, 24}; // T, 2N + N*N std::vector<int64_t> hPreShape = {1024, 4}; // T, N std::vector<int64_t> hPostShape = {1024, 4}; // T, N std::vector<int64_t> gammaShape = {4, 512}; // N, D std::vector<int64_t> gradXShape = {1024, 4, 512}; // T, N, D std::vector<int64_t> gradPhiShape = {24, 2048}; // 2N + N*N, N*D std::vector<int64_t> gradAlphaShape = {3}; std::vector<int64_t> gradGammaShape = {4, 512}; // N, D std::vector<int64_t> gradBiasShape = {24}; // 2N + N*N std::vector<int64_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::vector<short> xHostData(1024 * 4 * 512, 1.0); std::vector<float> phiHostData(24 * 2048, 1.0); std::vector<float> alphaHostData(3, 1.0); std::vector<short> gradHInHostData(1024 * 512, 1.0); std::vector<float> gradHPostHostData(1024 * 4, 1.0); std::vector<float> gradHResHostData(1024 * 4 * 4, 1.0); std::vector<float> invRmsHostData(1024, 1.0); std::vector<float> hMixHostData(1024 * 24, 1.0); std::vector<float> hPreHostData(1024 * 4, 1.0); std::vector<float> hPostHostData(1024 * 4, 1.0); std::vector<float> gammaHostData(4 * 512, 1.0); std::vector<short> gradXHostData(1024 * 4 * 512, 0); std::vector<float> gradPhiHostData(24 * 2048, 0); std::vector<float> gradAlphaHostData(3, 0); std::vector<float> gradBiasHostData(24, 0); std::vector<float> gradGammaHostData(4 * 512, 0); std::vector<short> 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 个输入xphialphagrad_h_ingrad_h_postgrad_h_resinv_rmsh_mixh_preh_post为必选,gammagrad_x_post为可选(OPTIONAL),与文档参数表一一对应;
  • 5 个输出grad_xgrad_phigrad_alphagrad_bias必选,grad_gamma可选;
  • 2 个属性hc_eps(FLOAT,默认1e-6f)与op_impl_mode(INT,默认 0);
  • 动态能力DynamicCompileStaticFlag(true)DynamicShapeSupportFlag(true)DynamicRankSupportFlag(true),支持动态 shape 与动态 rank(即 B、S 泛化)。

8.2 输出 shape 推导(infershape)

infershape 实现 的推导逻辑完全围绕"梯度 shape 必须与前向一致"展开:

  1. 校验维度组合ValidateInputDims要求gradHIngradHPost维数相等,且gradHRes必须匹配对应格式(BSND 时 gradHRes 为 4 维 BSNN 或 3 维 BSN!;TND 时 gradHRes 为 3 维 TNN 或 2 维 TN!);
  2. 推导 gradXgradX的 shape 由gradHIngradHPost组合得出,BS 格式下为 (B, S, N, D),T 格式下为 (T, N, D);
  3. 推导其余输出gradPhigradBias的第 0 维取自phi的第 0 维(fusionSize),gradAlpha固定为 3,gradGamma为 (N, D);
  4. 推导数据类型gradX继承grad_h_in的类型(BF16/FP16),其余梯度输出统一为 FLOAT32。

8.3 Tiling 分发与平台适配

Tiling 入口 mhc_pre_backward_tiling.cpp 根据 SoC 版本将计算任务分发给两套实现:

  • arch35(Ascend 950/A3/A5):通过 TilingRegistry 注册的模板实现,kernel 侧包含 Cube 计算(见 arch35 目录),支持 FP32/HF32 两种 Cube 模式,对应 V2 接口的opImplMode参数;
  • arch22(Ascend 910B/A2):独立的TilingMhcPreBackwardArch22实现(见 arch22 目录),对应 A2 平台仅支持 N=4、D 128 对齐的规格。

此外,TilingPrepare4MhcPreBackward在编译期采集各核 AIC/AIV 数量与 UB/L1/L2/L0 各级缓存大小,供 tiling 决策切分策略使用。

九、单元测试验证要点

算子的 op_api 单测位于 tests/ut/op_api/test_aclnn_mhc_pre_backward.cpp,主要覆盖以下行为:

  • 基础路径:N=4、D=8、T=10 的 TND 排布下,aclnnMhcPreBackward返回ACLNN_SUCCESS
  • 阶乘残差变体RunMhcPreBackward(true)使用fusionSize = N! + 2N(N=4 时为 24)的 shape,验证 A2 平台支持的 N! 形式融合维度;
  • V2 模式校验opImplMode=1(HF32)调用成功;opImplMode=2-1均返回失败,验证了"仅支持 0/1"的约束;
  • 空指针校验GetWorkspaceSize传入全空指针时返回ACLNN_ERR_PARAM_NULLPTR(161001);
  • 平台模拟:测试类通过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²+2N,A2 平台还支持 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/21 18:19:45

面试必问扎克加点:3个高频坑点与源码级避坑指南

面试必问扎克加点:3个高频坑点与源码级避坑指南 面试被问原理答不上来,这大概是每个程序员最尴尬的时刻。 尤其是当面试官盯着你的眼睛,追问“扎克加点”在并发场景下的具体表现时,你脑子里一片空白,只能支支吾吾地背八股文。 扎克加点 这个词,在常规技术栈里很少见,但在特定底层优化和高频交易系统中,它是…

作者头像 李华
网站建设 2026/9/21 18:19:39

亚洲中文字幕在线不卡电影实战项目避坑指南

亚洲中文字幕在线不卡电影实战项目避坑指南 版本升级后 API 全变了,这大概是很多做流媒体或视频处理的同学最头疼的事。昨天还在跑的代码,今天一升级依赖,直接报一堆找不到方法的错。在做一个 实战项目…

作者头像 李华
网站建设 2026/9/21 18:19:35

股市交易实战项目复盘:3个避坑技巧搞定面试原理

股市交易实战项目复盘:3个避坑技巧搞定面试原理 面试官问:“为什么你的量化策略在回测里赚钱,实盘就亏?”你愣住,答不上来。这种尴尬在技术圈太常见了。很多人把 股市交易 当成纯数学题,忽略了底层逻辑的落地。 别慌,今天不讲高深的金融理论,只讲怎么用编程思维搞定 股市交易 的 实战项目…

作者头像 李华
网站建设 2026/9/21 18:19:35

人生如逆旅我亦是行人项目避坑:3个最佳实践救活你的代码

人生如逆旅我亦是行人项目避坑:3个最佳实践救活你的代码 看了一堆教程还是不会写项目?别怪自己笨,是教程没讲透底层逻辑。很多人卡在“人生如逆旅我亦是行人”这种带有强业务含义或特定命名的模块里,死记硬背API却不懂数据流向。真正的 最佳实践 ,不是背代码,而是理解错误背后的机制。…

作者头像 李华
网站建设 2026/9/21 18:19:30

3个致命坑:如何分类汇总与源码解析避坑指南

3个致命坑:如何分类汇总与源码解析避坑指南 面对满屏红色的 Exception in thread "main" java.lang.NullPointerException ,你是不是也抓狂过?StackTrace 长得像天书,一行行看下去全是 at…

作者头像 李华
网站建设 2026/9/21 18:19:12

3个Rapier性能优化坑,面试原理秒答不慌

3个Rapier性能优化坑,面试原理秒答不慌 面试官问“物理引擎底层怎么保证稳定性”,你脑子一片空白?别慌。这不是你的错,是多数教程只教你调API,没讲透底层机制。今天拆解 Rapier 2D/3D 物理引擎在性能优化上的核心逻辑,帮你把“黑盒”变“白盒”。 概念速懂:为什么是 Rapier?…

作者头像 李华