残差连接后处理算子的原理、接口与调用实践)
算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载导读mhc_post是 CANN ops-transformer 仓库中 mHCManifold-Constraint Hyper-Connection流形约束超连接架构的核心后处理算子用于 Transformer 类大模型在昇腾 NPU 上完成多层残差连接的后处理阶段计算。它把残差矩阵变换Res Mapping与输出状态投影Post Mapping融合为一次调用避免多次独立算子带来的额外开销。阅读本文后你将掌握mhc_post的数学原理、完整参数约束、BSND/TND 两种维度格式的用法以及单算子模式与图模式torch.compile两种调用方式并能在自己的 PyTorch 工程中正确接入该算子。产品支持情况根据 mhc_post 算子说明 与 MhcPost README该算子在以下产品形态上的支持情况如下产品是否支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持需要特别注意的是h_res 是否可缺省在不同产品上有差异Ascend 950PR/Ascend 950DT 允许h_res传入None退化为直接残差连接而 Atlas A2/A3 系列产品要求h_res为必传参数传入None会直接报错。这一限制在算子定义与 ACLNN 入口中均有体现见下文源码分析。功能与数学原理接口功能mhc_post实现 MHC Post 组件的前向计算用于 Transformer 模型中多层残差连接的后处理阶段。该算子将残差矩阵变换Res Mapping与输出状态投影Post Mapping融合为单次计算对上一层输入 $x_l$ 使用转置后的残差矩阵 $H_l^{res}$ 做矩阵乘法变换对上一层输出 $h_l^{out}$ 使用后处理权重 $H_t^{post}$ 做逐元素缩放与广播二者相加得到下一层输入 $x_{l1}$。从 MhcPost README 的功能描述可以确认其定位MhcPost 基于一系列计算对 mHC 架构中上一层输出 $h_t^{out}$ 进行 Post Mapping对上一层的输入 $x_l$ 进行 Res Mapping然后对二者进行残差连接得到下一层的输入 $x_{l1}$。核心计算公式MHC Post 算子的核心计算公式为$$ x_{l1} (H_{l}^{res})^{T} \cdot x_{l} h_{l}^{out} \cdot H_{t}^{post} $$其中包含两个部分Res Mapping残差矩阵变换$(H_{l}^{res})^{T} \cdot x_{l}$ 表示对输入 $x_l$ 进行残差矩阵的转置矩阵乘法。对于输出中的第 $i$ 行对应第 $i$ 个 head计算过程为$$ x_{l1}[i] \sum_{j0}^{n-1} H_{l}^{res}[j, i] \cdot x_{l}[j] $$即将 $H_{l}^{res}$ 矩阵按转置方式与 $x_l$ 做矩阵乘法$H_{l}^{res}[j, i]$ 为标量对 $x_l$ 的第 $j$ 行做标量乘法后累加到第 $i$ 行输出。Post Mapping输出状态投影$h_{l}^{out} \cdot H_{t}^{post}$ 表示输出状态 $h_{l}^{out}$ 与后处理权重 $H_{t}^{post}$ 的逐元素乘法与广播。对于第 $i$ 个 head$$ x_{l1}[i] H_{t}^{post}[i] \cdot h_{l}^{out} $$即 $H_{t}^{post}[i]$ 为标量对 $h_{l}^{out}$ 整行做标量乘法后加到第 $i$ 行输出。综合完整计算过程为$$ x_{l1}[i, :] H_{t}^{post}[i] \cdot h_{l}^{out}[:] \sum_{j0}^{n-1} H_{l}^{res}[j, i] \cdot x_{l}[j, :] $$其中$x_{l}$ 对应参数x$H_{l}^{res}$ 对应参数h_res$h_{l}^{out}$ 对应参数h_out$H_{t}^{post}$ 对应参数h_post$x_{l1}$ 对应输出y。h_res 缺省时的退化形式当h_res传入 None 时跳过 Res Mapping计算公式退化为直接残差连接$$ x_{l1} x_{l} h_{l}^{out} \cdot H_{t}^{post} $$维度格式说明输入支持两种维度格式BSND4 维和TND3 维。其中BBatch批量大小SSeq-Length序列长度T所有 Batch 序列长度的累加和$T B \times S$n头数head 数量D每个头的隐藏维度大小headdim。源码级印证算子定义见 mhc_post_def.cppx/h_out声明为REQUIRED且支持ge::DT_FLOAT16与ge::DT_BF16h_res/h_post声明为OPTIONAL/REQUIRED且仅支持ge::DT_FLOATfloat32格式均为FORMAT_ND同时为ascend910b、ascend910_93对应 A2/A3以及ascend950、ascend350配置了 AICore 实现。ACLNN 入口见 aclnn_mhc_post.cpp在aclnnMhcPostGetWorkspaceSize中通过op::GetCurrentPlatformInfo().GetCurNpuArch() ! NpuArch::DAV_3510对h_res nullptrnohres 路径做了 SoC 守卫仅在 Ascend 950 架构DAV_3510放行与文档仅 Ascend 950PR/Ascend 950DT 支持传入 None的描述一致。PyTorch 封装见 mhc_post.py通过torch.ops.cann_ops_transformer.mhc_post完成算子调度并提供torch.autograd.Function封装以支持反向传播。Golden 参考实现见 golden.pygolden 与 third_party 均按文档公式用 torch 小算子拼接验证nohres 变体对应y x h_post.unsqueeze(-1) * h_out.unsqueeze(-2)所有计算在 float32 下完成后 cast 回x的 dtype可作为理解计算语义的最佳参考。函数原型cann_ops_transformer.mhc_post(x, h_res, h_out, h_post) - Tensor其中h_res为可选输入可传入 None仅 Ascend 950PR/Ascend 950DT 的单算子模式支持传入 None图模式不支持。从 mhc_post.py 中的算子 schema 可以确认底层签名mhc_post(Tensor x, Tensor? hRes, Tensor hOut, Tensor hPost) - Tensor其中hRes声明为可空Tensor?与 Python 层h_res可传None的语义一致。参数说明下表为mhc_post的完整参数说明参数名参数类型可选/必选描述数据类型维度(shape)xTensor必选当前层的输入 token 特征对应公式中的 $x_l$。bfloat16、float16(B, S, n, D) 或 (T, n, D)h_resTensor可选残差连接矩阵对应公式中的 $H_l^{res}$。传入 None 时跳过 Res Mapping计算公式退化为直接残差连接仅 Ascend 950PR/Ascend 950DT 支持传入 None其他产品形态传入 None 会报错。float32(B, S, n, n) 或 (T, n, n)h_outTensor必选上一层的输出状态对应公式中的 $h_l^{out}$。bfloat16、float16(B, S, D) 或 (T, D)h_postTensor必选后处理权重矩阵对应公式中的 $H_t^{post}$。float32(B, S, n) 或 (T, n)参数语义补充h_res 的物理含义从 MhcPost README 可知h_res是 mHC 的 h_res 变换矩阵是做完 Sinkhorn 变换后的双随机矩阵h_out是 Atten/MLP 层的输出h_post是 mHC 的 h_post 变换矩阵。结合仓库中 mhc_pre_sinkhorn 等前向组件可以推断h_res由 mHC 前级如 sinkhorn 归一化产生mhc_post负责在层间残差连接处消费这些矩阵。与 README 的规格差异说明README 中标注了n 固定为 4、d 为 128 的倍数范围 1 到 100000的参考规格而 torch API 文档本文主体未限制 n 的具体取值仅要求各维度为正数。以 torch API 文档为准的同时若参考 README 的典型配置n4、D128可复现官方测试用例的常见形态如 test_mhc_post_infershape.cpp 中的 4D 用例 shape 为 (512, 2, 4, 512)、3D 用例为 (1024, 4, 512)。返回值说明参数名参数类型可选/必选描述数据类型维度(shape)yTensor必选MHC Post 计算输出对应公式中的 $x_{l1}$数据类型与输入x保持一致shape 与输入x保持一致。bfloat16、float16(B, S, n, D) 或 (T, n, D)该语义在 mhc_post_infershape.cpp 中得到印证infer shape 将输出y的维度数设置为与x相同并逐维拷贝yShape-SetDim(i, xShape-GetDim(i))infer data type 则将输出类型直接设为x的类型context-SetOutputDataType(INDEX_Y, xDtype)。此外未知 rankIsUnknownRank场景下输出 shape 也被置为未知 rank支持动态形状传播。约束说明使用mhc_post时需满足以下约束使用场景该接口支持训练、推理场景下使用。调用模式该接口支持单算子模式和图模式调用。数据类型约束x和h_out的数据类型必须相同输出y的数据类型与x保持一致。h_res 可选约束h_res传入 None 时跳过 Res Mapping仅 Ascend 950PR/Ascend 950DT 支持Atlas A2/A3 系列产品h_res为必传参数传入 None 会报错h_res传入 None 仅支持单算子模式调用图模式torch.compile下h_res必须传入传入 None 会在 GE 编译阶段报错。维度约束h_res的维度需与x维度格式匹配4 维时为 (B, S, n, n)3 维时为 (T, n, n)。Shape 一致性约束4 维BSND格式下h_res的 (B, S) 维度需与x的 (B, S) 维度一致h_res的后两维为 (n, n)其中 n 与x的第 3 维一致h_out的 (B, S) 维度需与x的 (B, S) 维度一致h_out的 D 维度需与x的 D 维度一致h_post的 (B, S) 维度需与x的 (B, S) 维度一致h_post的 n 维度需与x的 n 维度一致。3 维TND格式下h_res的 T 维度需与x的 T 维度一致后两维为 (n, n)h_out的 T 维度需与x的 T 维度一致D 维度需与x的 D 维度一致h_post的 T 维度需与x的 T 维度一致n 维度需与x的 n 维度一致。正数约束所有输入 Tensor 的 shape 各维度值必须为正数大于 0。这些 Shape 一致性检查在 ACLNN 层的 aclnn_mhc_post.cpp 中有完整的运行时校验实现CheckShape3D/CheckShape4D逐一比对hRes、hOut、hPost、out各维度与x的对应维度例如 4D 下要求hResDim3 hResDim2n×n 矩阵、hOutDim2 xDim3、hPostDim2 xDim2CheckDtype校验x/hOut/out类型一致、hRes/hPost必须为 FP32。因此传入不匹配的 shape 或 dtype 时调用会返回ACLNN_ERR_PARAM_INVALID而非静默出错。确定性计算默认支持确定性计算。调用说明单算子模式调用以下示例展示在昇腾 NPU 上以单算子模式调用mhc_post完整代码可参考 examples/test_aclnn_mhc_post.cpp 的 C 对应实现以及 torch_extension 封装 mhc_post.pyimport torch import torch_npu from cann_ops_transformer.ops import mhc_post B 2 S 8 n 4 D 128 x torch.randn(B, S, n, D, dtypetorch.bfloat16).npu() h_res torch.randn(B, S, n, n, dtypetorch.float32).npu() h_out torch.randn(B, S, D, dtypetorch.bfloat16).npu() h_post torch.randn(B, S, n, dtypetorch.float32).npu() y mhc_post(x, h_res, h_out, h_post) print(foutput shape: {y.shape}) # h_res缺省时仅Ascend 950PR/Ascend 950DT支持且仅支持单算子模式 y mhc_post(x, None, h_out, h_post) print(foutput shape: {y.shape})代码要点输入x、h_out使用bfloat16也可用float16h_res、h_post使用float32BSND 4 维格式下x为 (B, S, n, D)h_res为 (B, S, n, n)h_out为 (B, S, D)h_post为 (B, S, n)输出y的 shape 与 dtype 均与x一致在非 Ascend 950 产品上传入h_resNone会报错在 Ascend 950 上也需要单算子模式才能使用该缺省路径。如需在 BSND 与 TND 格式间切换将 4 维输入合并为 3 维即可把x视为 (T, n, D)h_res视为 (T, n, n)h_out视为 (T, D)h_post视为 (T, n)其中 $T B \times S$如 test_mhc_post_infershape.cpp 中的 3D 用例 (1024, 4, 512)。图模式调用torch.compile图模式通过 torchair 的 NPU 后端将mhc_post转换为 GE 图节点MhcPost执行。注意图模式下h_res必须传入不能为 Noneimport torch import torch_npu import torchair from cann_ops_transformer.ops import mhc_post torch_npu.npu.set_device(0) B 2 S 8 n 4 D 128 class MhcPostModel(torch.nn.Module): def forward(self, x, h_res, h_out, h_post): return mhc_post(x, h_res, h_out, h_post) model MhcPostModel().npu() npu_backend torchair.get_npu_backend() model torch.compile(model, backendnpu_backend, dynamicFalse) x torch.randn(B, S, n, D, dtypetorch.bfloat16, devicenpu) h_res torch.randn(B, S, n, n, dtypetorch.float32, devicenpu) h_out torch.randn(B, S, D, dtypetorch.bfloat16, devicenpu) h_post torch.randn(B, S, n, dtypetorch.float32, devicenpu) y model(x, h_res, h_out, h_post)图模式的实现原理仓库中的 graph_convert_mhc_post.py 实现了 FX 节点到 GE 算子的转换器。它通过register_fx_node_ge_converter(torch.ops.cann_ops_transformer.mhc_post.default)注册转换函数将torch.ops.cann_ops_transformer.mhc_post转换为 GE 图中的MhcPost算子节点其 IR 定义为输入xDT_BF16 / DT_FLOAT16、h_resDT_FLOAT、h_outDT_BF16 / DT_FLOAT16、h_postDT_FLOAT输出yDT_BF16 / DT_FLOAT16。这也从侧面说明为什么图模式下h_res必须传入GE 图中MhcPost节点的h_res输入在 converter 中始终作为必填张量下发无法表达缺省语义。反向传播与自动微分mhc_post的 PyTorch 封装 mhc_post.py 提供了自动微分支持当任一输入x、h_out、h_post以及非 None 的h_res需要梯度时走MhcPostFunctiontorch.autograd.Function路径backward通过mhc_post_backward算子计算梯度并针对h_resNone场景将grad_h_res置为 None为规避 0-stride 的 expanded grad_output例如sum().backward()产生在h_res为 None 时破坏aclnnMhcPostBackward的问题反向入口先将grad_output.contiguous()物化。这一设计与仓库中的 mhc_post_backward 组件配合使用说明mhc_post在训练场景下不仅支持前向计算还具备完整的梯度链路。总结mhc_post将 mHC 架构中残差矩阵转置乘法 输出状态投影 残差连接三步融合为单算子同时支持 BSND4 维与 TND3 维两种布局覆盖训练与推理、单算子与图模式四大使用场景。使用时的关键要点可归纳为dtype 搭配x/h_out同为 bf16 或 fp16h_res/h_post必须为 fp32输出与x同 dtypeh_res 缺省仅 Ascend 950PR/Ascend 950DT 单算子模式支持 None其余产品及图模式必须显式传入shape 一致性h_res的 (n, n)、h_out的 D、h_post的 n 必须与x严格对应ACLNN 层会做完整运行时校验确定性算子默认支持确定性计算便于调试与结果复现。如需进一步深入可继续阅读 aclnnMhcPost 接口文档C/ACLNN 调用方式、MhcPost README算子级规格与调用方式汇总、op_host 实现算子定义与 tiling 逻辑以及 tests 目录infershape 与 tiling 单元测试、golden 参考实现。赞分享算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载相关推荐CANN ops-transformer 算子 aclnnMhcPost 使用指南mHC 架构 Post Mapping 与残差连接的 NPU 融合实现CANN ops transformer 算子 aclnnMhcPost 使用指南mHC 架构 Post Mapping 与残差连接的 NPU 融合实现 导读算子库人工智能深度学习Ascendmhc_post 算子实战解析CANN ops-transformer 中 mHC 后连接的广播缩放 AscendC 实现mhc_post 算子实战解析CANN ops transformer 中 mHC 后连接的广播缩放 AscendC 实现 本文围绕 CANN ops tra算子库人工智能深度学习AscendCANN ops-transformer 中的 mHC 流形约束超连接 AscendC 算子mhc_pre / mhc_post / mhc_res 实现与实战指南CANN ops transformer 中的 mHC 流形约束超连接 AscendC 算子mhc_pre / mhc_post / mhc_res 实现与实算子库人工智能深度学习Ascend上一篇一文搞懂 AssetRipper把 Unity 黑盒资源拆成你能编辑的文件下一篇优化Renovate版本转换从混乱到自动化的依赖管理革命创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考