CANN ops-transformer 算子指南:MHC Post(mhc_post)残差连接后处理算子的原理、接口与调用实践
2026/9/20 9:10:37 网站建设 项目流程
  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

导读

mhc_post是 CANN ops-transformer 仓库中 mHC(Manifold-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_{l+1}$。

从 MhcPost README 的功能描述可以确认其定位:MhcPost 基于一系列计算对 mHC 架构中上一层输出 $h_t^{out}$ 进行 Post Mapping,对上一层的输入 $x_l$ 进行 Res Mapping,然后对二者进行残差连接,得到下一层的输入 $x_{l+1}$。

核心计算公式

MHC Post 算子的核心计算公式为:

$$ x_{l+1} = (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_{l+1}[i] = \sum_{j=0}^{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_{l+1}[i] += H_{t}^{post}[i] \cdot h_{l}^{out} $$

即 $H_{t}^{post}[i]$ 为标量,对 $h_{l}^{out}$ 整行做标量乘法后加到第 $i$ 行输出。

综合完整计算过程为:

$$ x_{l+1}[i, :] = H_{t}^{post}[i] \cdot h_{l}^{out}[:] + \sum_{j=0}^{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_{l+1}$ 对应输出y

h_res 缺省时的退化形式:当h_res传入 None 时,跳过 Res Mapping,计算公式退化为直接残差连接:

$$ x_{l+1} = x_{l} + h_{l}^{out} \cdot H_{t}^{post} $$

维度格式说明

输入支持两种维度格式:BSND(4 维)TND(3 维)。其中:

  • B(Batch):批量大小;
  • S(Seq-Length):序列长度;
  • T:所有 Batch 序列长度的累加和,$T = B \times S$;
  • n:头数(head 数量);
  • D:每个头的隐藏维度大小(headdim)。

源码级印证

  • 算子定义见 mhc_post_def.cpp:x/h_out声明为REQUIRED且支持ge::DT_FLOAT16ge::DT_BF16h_res/h_post声明为OPTIONAL/REQUIRED且仅支持ge::DT_FLOAT(float32),格式均为FORMAT_ND;同时为ascend910bascend910_93(对应 A2/A3)以及ascend950ascend350配置了 AICore 实现。
  • ACLNN 入口见 aclnn_mhc_post.cpp:在aclnnMhcPostGetWorkspaceSize中通过op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510h_res == nullptr(nohres 路径)做了 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.py:golden 与 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 的典型配置(n=4、D=128)可复现官方测试用例的常见形态(如 test_mhc_post_infershape.cpp 中的 4D 用例 shape 为 (512, 2, 4, 512)、3D 用例为 (1024, 4, 512))。

返回值说明

参数名参数类型可选/必选描述数据类型维度(shape)
yTensor必选MHC Post 计算输出,对应公式中的 $x_{l+1}$,数据类型与输入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))。此外,未知 rank(IsUnknownRank)场景下输出 shape 也被置为未知 rank,支持动态形状传播。

约束说明

使用mhc_post时需满足以下约束:

  • 使用场景:该接口支持训练、推理场景下使用。
  • 调用模式:该接口支持单算子模式和图模式调用。
  • 数据类型约束
    • xh_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逐一比对hReshOuthPostout各维度与x的对应维度(例如 4D 下要求hResDim3 == hResDim2(n×n 矩阵)、hOutDim2 == xDim3hPostDim2 == xDim2),CheckDtype校验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.py):

import 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, dtype=torch.bfloat16).npu() h_res = torch.randn(B, S, n, n, dtype=torch.float32).npu() h_out = torch.randn(B, S, D, dtype=torch.bfloat16).npu() h_post = torch.randn(B, S, n, dtype=torch.float32).npu() y = mhc_post(x, h_res, h_out, h_post) print(f"output shape: {y.shape}") # h_res缺省时(仅Ascend 950PR/Ascend 950DT支持,且仅支持单算子模式) y = mhc_post(x, None, h_out, h_post) print(f"output shape: {y.shape}")

代码要点:

  • 输入xh_out使用bfloat16(也可用float16),h_resh_post使用float32
  • BSND 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_res=None会报错;在 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必须传入,不能为 None:

import 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, backend=npu_backend, dynamic=False) x = torch.randn(B, S, n, D, dtype=torch.bfloat16, device="npu") h_res = torch.randn(B, S, n, n, dtype=torch.float32, device="npu") h_out = torch.randn(B, S, D, dtype=torch.bfloat16, device="npu") h_post = torch.randn(B, S, n, dtype=torch.float32, device="npu") 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 定义为:

  • 输入x(DT_BF16 / DT_FLOAT16)、h_res(DT_FLOAT)、h_out(DT_BF16 / DT_FLOAT16)、h_post(DT_FLOAT);
  • 输出y(DT_BF16 / DT_FLOAT16)。

这也从侧面说明为什么图模式下h_res必须传入:GE 图中MhcPost节点的h_res输入在 converter 中始终作为必填张量下发,无法表达"缺省"语义。

反向传播与自动微分

mhc_post的 PyTorch 封装 mhc_post.py 提供了自动微分支持:

  • 当任一输入(xh_outh_post,以及非 None 的h_res)需要梯度时,走MhcPostFunctiontorch.autograd.Function)路径;
  • backward通过mhc_post_backward算子计算梯度,并针对h_res=None场景将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 架构中"残差矩阵转置乘法 + 输出状态投影 + 残差连接"三步融合为单算子,同时支持 BSND(4 维)与 TND(3 维)两种布局,覆盖训练与推理、单算子与图模式四大使用场景。使用时的关键要点可归纳为:

  1. dtype 搭配x/h_out同为 bf16 或 fp16,h_res/h_post必须为 fp32,输出与x同 dtype;
  2. h_res 缺省:仅 Ascend 950PR/Ascend 950DT 单算子模式支持 None,其余产品及图模式必须显式传入;
  3. shape 一致性h_res的 (n, n)、h_out的 D、h_post的 n 必须与x严格对应,ACLNN 层会做完整运行时校验;
  4. 确定性:算子默认支持确定性计算,便于调试与结果复现。

如需进一步深入,可继续阅读 aclnnMhcPost 接口文档(C++/ACLNN 调用方式)、MhcPost README(算子级规格与调用方式汇总)、op_host 实现(算子定义与 tiling 逻辑)以及 tests 目录(infershape 与 tiling 单元测试、golden 参考实现)。

  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

相关推荐

上一篇:一文搞懂 AssetRipper:把 Unity 黑盒资源拆成你能编辑的文件
下一篇:优化Renovate版本转换:从混乱到自动化的依赖管理革命

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询