- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
导读
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_FLOAT16与ge::DT_BF16,h_res/h_post声明为OPTIONAL/REQUIRED且仅支持ge::DT_FLOAT(float32),格式均为FORMAT_ND;同时为ascend910b、ascend910_93(对应 A2/A3)以及ascend950、ascend350配置了 AICore 实现。 - ACLNN 入口见 aclnn_mhc_post.cpp:在
aclnnMhcPostGetWorkspaceSize中通过op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510对h_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) |
|---|---|---|---|---|---|
| x | Tensor | 必选 | 当前层的输入 token 特征,对应公式中的 $x_l$。 | bfloat16、float16 | (B, S, n, D) 或 (T, n, D) |
| h_res | Tensor | 可选 | 残差连接矩阵,对应公式中的 $H_l^{res}$。传入 None 时跳过 Res Mapping,计算公式退化为直接残差连接,仅 Ascend 950PR/Ascend 950DT 支持传入 None,其他产品形态传入 None 会报错。 | float32 | (B, S, n, n) 或 (T, n, n) |
| h_out | Tensor | 必选 | 上一层的输出状态,对应公式中的 $h_l^{out}$。 | bfloat16、float16 | (B, S, D) 或 (T, D) |
| h_post | Tensor | 必选 | 后处理权重矩阵,对应公式中的 $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) |
|---|---|---|---|---|---|
| y | Tensor | 必选 | 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时需满足以下约束:
- 使用场景:该接口支持训练、推理场景下使用。
- 调用模式:该接口支持单算子模式和图模式调用。
- 数据类型约束:
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 维度一致。
- 4 维(BSND)格式下:
- 正数约束:所有输入 Tensor 的 shape 各维度值必须为正数(大于 0)。
这些 Shape 一致性检查在 ACLNN 层的 aclnn_mhc_post.cpp 中有完整的运行时校验实现:CheckShape3D/CheckShape4D逐一比对hRes、hOut、hPost、out各维度与x的对应维度(例如 4D 下要求hResDim3 == hResDim2(n×n 矩阵)、hOutDim2 == xDim3、hPostDim2 == 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}")代码要点:
- 输入
x、h_out使用bfloat16(也可用float16),h_res、h_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 提供了自动微分支持:
- 当任一输入(
x、h_out、h_post,以及非 None 的h_res)需要梯度时,走MhcPostFunction(torch.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 维)两种布局,覆盖训练与推理、单算子与图模式四大使用场景。使用时的关键要点可归纳为:
- dtype 搭配:
x/h_out同为 bf16 或 fp16,h_res/h_post必须为 fp32,输出与x同 dtype; - h_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上加速计算。
相关推荐
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
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考