CANN ops-transformer 算子解析:MhcPost(mHC 架构 Post/Res Mapping 与残差融合算子)原理与调用实战
2026/9/20 18:16:14 网站建设 项目流程
  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

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

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

导读

MhcPost 是 CANN ops-transformer 算子库中 mHC(Multi-Head Collaboration,多头协同)架构的关键后处理算子,它将 mHC 架构上一层的输出h_out做 Post Mapping、上一层的输入x做 Res Mapping 后执行残差连接,一步融合生成下一层输入x_{l+1}。本文基于 mhc/mhc_post/README.md 及配套的 aclnnMhcPost 接口文档、PyTorch API 文档,结合仓库源码与测试用例,完整讲解其计算原理、算子规格、两种调用方式(aclnn C++ 接口与 PyTorch API)以及底层 tiling/vector 实现,帮助开发者在 Ascend NPU 上正确配置与调用该算子。

MhcPost 在 mHC 架构中的定位

mHC(Multi-Head Collaboration)是面向 Transformer 大模型的多头协同计算架构。在 mHC 中,每一层 Transformer 块的前后都会插入"映射"计算,用于在多头注意力/MLP 子层与层间残差流之间做信息变换。MhcPost 正是承担这一"层后处理"职责的算子:它位于某一层(Atten/MLP 子层)之后,将该层输出与上一层输入重新组织成下一层的输入,从而将"后处理投影 + 残差映射 + 残差相加"三类计算融合为单次算子调用,避免多次独立算子调用带来的额外搬运与调度开销。

从 PyTorch API 文档 的功能描述可以确认其定位:"实现 MHC Post 组件的前向计算,用于 Transformer 模型中多层残差连接的后处理阶段。该算子将残差矩阵变换与输出状态投影融合为单次计算,避免多次独立算子调用带来的额外开销。"

计算原理与公式

完整计算(h_res 提供时)

设上一层输入为 $x_l$,上一层输出(Atten/MLP 层输出)为 $h_{l}^{out}$,mHC 的残差映射矩阵为 $H_{l}^{res}$(sinkhorn 变换后的双随机矩阵),后处理映射矩阵为 $H_{t}^{post}$,则下一层输入为:

$$ x_{l+1} = (H_{l}^{res})^{T} \times x_l + h_{l}^{out} \otimes H_{t}^{post} $$

其中 $\otimes$ 表示逐元素乘法与广播,两部分语义如下(对应 torchapi_mhc_post.md 中的逐 head 展开式):

  • Res Mapping(对输入 $x_l$ 做残差矩阵转置乘法):对输出的第 $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}$ 做后处理投影):对第 $i$ 个 head:

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

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

综合逐 head 完整形式为:

$$ 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/out

退化路径(h_res 缺省时)

h_res传入nullptr/None时,跳过 Res Mapping,公式退化为直接残差连接:

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

该退化路径仅 Ascend 950PR/Ascend 950DT 支持(见下文"产品支持与约束")。

数据布局约定

算子支持两种输入维度格式(详见 torchapi_mhc_post.md 的维度说明):

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

产品支持情况与约束说明

产品支持矩阵

依据 README.md 与 aclnnMhcPost.md,支持情况如下:

产品是否支持
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支持传入nullptr,此时退化为直接残差连接;
    • Atlas A2/A3 训练与推理系列产品:h_res为必传参数,不支持传入nullptr(传入会报错)。
  • 规格约束

    规格项规格规格说明
    n4固定为 4
    d范围 1 到 100000128 的倍数
  • 确定性计算aclnnMhcPost默认确定性实现(见 aclnnMhcPost.md 约束说明)。

  • 数据类型约束(来自 torchapi_mhc_post.md):xh_out数据类型必须相同;输出y数据类型与x保持一致;h_resh_post为 FLOAT32。

  • Shape 一致性:以 BSND 格式为例,h_res(B, S)需与x一致、后两维为(n, n)h_out(B, S)需与x一致、D 维与x的 D 维一致;h_post(B, S)x一致、n 维与x的 n 维一致。所有输入 Tensor 各维度值必须为正数(大于 0)。

  • 图模式限制h_res传入None仅支持单算子模式调用;图模式(torch.compile)下h_res必须传入,否则会在 GE 编译阶段报错。

参数说明(算子输入/输出)

依据 README.md 的参数表,算子输入输出定义如下:

参数名输入/输出描述数据类型数据格式
x输入待计算的张量,表示网络中 mHC 层的输入数据FLOAT16、BFLOAT16ND
h_res输入(可选)mHC 的 h_res 变换矩阵,是做完 sinkhorn 变换后的双随机矩阵。缺省时退化为直接残差连接(仅 Ascend 950 支持)FLOAT32ND
h_out输入Atten/MLP 层的输出FLOAT16、BFLOAT16ND
h_post输入mHC 的 h_post 变换矩阵FLOAT32ND
out输出网络中 mHC 层的输出数据,作为下一层的输入FLOAT16、BFLOAT16ND

从算子注册源码 mhc/mhc_post/op_host/mhc_post_def.cpp 可以看到,x/h_out为 REQUIRED 且数据类型为DT_FLOAT16/DT_BF16h_res为 OPTIONAL 且为DT_FLOATh_post为 REQUIRED 且为DT_FLOAT,输出yx同类型;同时配置了ascend910bascend910_93ascend950ascend350四类 AICore 配置,其中 950/350 走mhc_post_apt实现(见ExtendCfgInfo("opFile.value", "mhc_post_apt")),910b/910_93 走mhc_post实现。

调用方式一:aclnn C++ 接口

两段式接口原型

MhcPost 的 aclnn 调用采用 CANN 标准两段式接口:先调用aclnnMhcPostGetWorkspaceSize获取 workspace 大小与执行器,再调用aclnnMhcPost执行计算。

aclnnStatus aclnnMhcPostGetWorkspaceSize( const aclTensor *x, const aclTensor *hRes, // 可选,Ascend 950 上可传 nullptr const aclTensor *hOut, const aclTensor *hPost, aclTensor *out, uint64_t *workspaceSize, aclOpExecutor **executor)
aclnnStatus aclnnMhcPost( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)

GetWorkspaceSize 参数与返回码

aclnnMhcPostGetWorkspaceSize中各张量参数(详见 aclnnMhcPost.md)的 shape 与约束如下:

参数输入/输出数据类型维度(shape)非连续 Tensor
x输入FLOAT16、BFLOAT16[B,S,N,D]、[T,N,D]
hRes输入(可选)FLOAT32[B,S,N,N]、[T,N,N]
hOut输入与 x 相同[B,S,D]、[T,D]
hPost输入FLOAT32[B,S,N]、[T,N]
out输出与 x 相同[B,S,N,D]、[T,N,D]-
workspaceSize输出---
executor输出---

第一段接口完成入参校验,常见返回码如下(完整返回码语义参见 aclnn 返回码说明):

返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001x、hOut、hPost、out 存在空指针
ACLNN_ERR_PARAM_INVALID161002x、hRes(非空时)、hOut、hPost、out 的数据类型不在支持范围内;或 shape 维度不在支持范围内;或数据类型/shape 不匹配
ACLNN_ERR_INNER_NULLPTR561103n 不等于 4;d 不在 [1, 100000] 范围内;d 不能被 128 整除

需要特别说明的是,hRes == nullptr的 SoC 校验在源码 aclnn_mhc_post.cpp 的aclnnMhcPostGetWorkspaceSize中显式实现:当hResnullptr且当前 NPU 架构不是DAV_3510(即 Ascend 950)时,直接返回ACLNN_ERR_RUNTIME_ERROR。这是"h_res 缺省仅 Ascend 950 支持"这一约束在接口层的落地保证。

完整调用示例

仓库提供了可直接参考的样例 mhc/mhc_post/examples/test_aclnn_mhc_post.cpp(README 中链接为examples/test_aclnn_mhc_post.cpp),完整的调用流程与 aclnnMhcPost.md 中的示例代码一致,核心步骤:

  1. 设备与流初始化aclInitaclrtSetDevice(deviceId)aclrtCreateStream
  2. 构造输入输出 Tensor:以 BSND 格式为例,典型 shape 为x = {1, 1024, 4, 5120}(BSND)、hRes = {1, 1024, 4, 4}hOut = {1, 1024, 5120}(BSD)、hPost = {1, 1024, 4}(BSn)、out = {1, 1024, 4, 5120}。通过aclrtMalloc申请 Device 内存、aclrtMemcpy将 Host 数据拷入,再用aclCreateTensorACL_FORMAT_ND创建aclTensor(x/hOut 用ACL_FLOAT16,hRes/hPost 用ACL_FLOAT);
  3. 调用第一段接口获取 workspace 大小与执行器,若workspaceSize > 0aclrtMalloc申请 workspace 内存;
  4. 调用第二段接口aclnnMhcPost(workspaceAddr, workspaceSize, executor, stream)执行计算;
  5. 同步并取回结果aclrtSynchronizeStreamaclrtMemcpy(DEVICE_TO_HOST)拷贝输出;
  6. 释放资源aclDestroyTensoraclrtFreeaclrtDestroyStreamaclrtResetDeviceaclFinalize

从接口实现源码 aclnn_mhc_post.cpp 可以确认两段式接口内部的完整执行链:第一段接口依次完成格式校验(CheckFormat,拒绝私有格式)、数据类型校验(CheckDtype,x/hOut 限定 FP16/BF16 且同型,hRes/hPost 限定 FP32)、Shape 校验(CheckShape,区分 3D TND 与 4D BSND 两套约束)以及空 Tensor 短路处理;随后通过l0op::Contiguous将非连续输入转连续(README 中标注"非连续 Tensor √"即由此支持),再调用底层l0op::MhcPost构建算子执行器,最后用l0op::ViewCopy将结果写回输出并返回 workspace 大小。第二段接口则通过CommonOpExecutorRun提交 AICore 任务。

调用方式二:PyTorch API

函数原型

cann_ops_transformer.mhc_post(x, h_res, h_out, h_post) -> Tensor

其中h_res为可选输入,可传入None(仅 Ascend 950PR/Ascend 950DT 的单算子模式支持传入None,图模式不支持)。

参数与返回值

依据 torchapi_mhc_post.md:

参数名参数类型可选/必选描述数据类型维度(shape)
xTensor必选当前层的输入 token 特征,对应公式中的 $x_l$bfloat16、float16(B, S, n, D)、(T, n, D)
h_resTensor可选残差连接矩阵,对应 $H_{l}^{res}$;传入 None 时跳过 Res Mappingfloat32(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)

返回值y:MHC Post 计算输出,对应公式中的 $x_{l+1}$,数据类型与x一致,shape 与x一致((B, S, n, D) 或 (T, n, D))。

单算子模式调用示例

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}")

图模式调用示例

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)

Torch 扩展实现要点

从 torch_extension/mhc_post.py 源码可以看到 PyTorch 侧的封装机制:

  • 通过OpBuilder动态编译csrc/mhc/mhc_post.cpp并注册自定义算子mhc_post(Tensor x, Tensor? hRes, Tensor hOut, Tensor hPost) -> Tensor(注意 schema 中hRes声明为可空Tensor?);
  • 当任一输入需要梯度时,自动走torch.autograd.FunctionMhcPostFunction),其backward会调用仓库中配套的mhc_post_backward(mhc/mhc_post_backward)计算四个输入的梯度;h_res is None时梯度返回None
  • 反向实现中还将grad_outputcontiguous()物化,以规避 h_res 为 None 时 0-stride 的扩展梯度(如sum().backward()产生)触发aclnnMhcPostBackward异常的问题。

底层实现:Tiling 与 Vector 计算

算子注册与 Shape 推导

  • 算子定义位于 mhc_post_def.cpp,通过OpDef注册输入输出、数据类型与各 SoC 的 AICore 配置;
  • Shape/数据类型推导位于 mhc_post_infershape.cpp:输出y的 shape 与x逐维一致,数据类型与x相同;x的维度必须为 3(TND)或 4(BSND),否则推导失败。

Tiling 机制

mhc_post_tiling_base.cpp 通过TilingRegistryArch按架构分发 tiling 实现:

  • arch22(对应 Atlas A2 系列,ascend910b):mhc_post_tiling_base_arch22.cpp;
  • arch35(对应 Ascend 950/A3 系列,ascend950/ascend350):分为三套实现 mhc_post_tiling_base_arch35.cpp(常规路径)、mhc_post_tiling_nohres_arch35.cpp(h_res 缺省退化路径)与 mhc_post_tiling_regbase_arch35.cpp(regbase 路径)。

Tiling 的核心任务是把(BS, n, d)维度的计算按"d 方向切块(dInner/dTail)"和"多核均分(normalCoreProcessNum/tailCoreProcessNum)"展开,从 kernel 代码中的bsIdx = globalItemIdx / dOuterdIdx = globalItemIdx % dOuter可看出,每个计算 item 是一个(bs, d 块)组合。

Vector Kernel 实现要点

主 kernel 入口 mhc_post.cpp 通过模板参数usePermanentX区分两种路径,计算实现集中在 arch22/mhc_post_arch22.h,关键优化点包括:

  • Double Buffer 流水:输入队列hOutTileQueue_xTileQueue_与输出队列outputTileQueue_深度均为 2,注释明确说明"Double Buffer 提升 Memory Bound 算子性能";
  • FP32 中间计算hOut/x先经Cast转为 F32(hOutF32Buf_xF32Buf_),Post Mapping 用Muls(outF32, hOutF32, hPost[i], dNum)实现标量乘,Res Mapping 用Axpy(outF32, xF32, hRes[j*n+i], dNum)实现转置矩阵乘累加(下标j * n + i正是 $H^{res}$ 转置访问的体现),最后Cast(CAST_RINT)回合回 FP16/BF16 输出;
  • usePermanentX 优化:当USE_PERMANENT_X == 1时一次性把 n 行x全部搬入(DataCopyExtParams多行 stride 拷贝),内层循环只做Axpy累加,避免反复搬 x 数据。

h_res 缺省路径则由 arch35/mhc_post_nohres.h 实现(MhcPostNoHRes类,公式x_{l+1} = x_l + h_{l}^{out} * H_{t}^{post}),同样采用 Double Buffer 队列(hOutTileQueue_/xTileQueue_/hPostTileQueue_)与DoMulAndAdd完成乘加融合。

单元测试

仓库为算子提供了较完整的 UT 覆盖(mhc/mhc_post/tests/ut):

  • Shape 推导测试:test_mhc_post_infershape.cpp;
  • 各架构 Tiling 测试:arch22/test_mhc_post_tiling.cpp、arch35/test_mhc_post_tiling.cpp;
  • aclnn 接口测试:op_api/test_aclnn_mhc_post.cpp。

小结与选型建议

MhcPost 将 mHC 架构层间后处理所需的 Post Mapping、Res Mapping 与残差连接融合为单算子,既可用于 aclnn 两段式 C++ 编程(调用样例),也可通过cann_ops_transformer.mhc_post在 PyTorch 单算子/图模式下直接调用。开发时需重点把握以下几点:

  1. 平台差异h_res缺省的退化路径仅 Ascend 950PR/Ascend 950DT 支持(接口层由DAV_3510架构校验保证);Atlas A2/A3 产品必须显式传入h_res
  2. 规格合规:n 固定为 4,d 必须在 [1, 100000] 且为 128 的倍数,否则接口层返回 561103 错误;
  3. 类型与 Shapex/h_out必须同为 FP16 或 BF16,h_res/h_post为 FP32;BSND 与 TND 两套 shape 需按上文约束严格对齐,输出与x同型同 shape;
  4. 模式差异h_res=None仅支持单算子模式;图模式(torch.compile)下必须传入h_res,否则 GE 编译阶段报错。

对需要反向传播的场景,可直接使用 PyTorch 扩展层自动接入 mhc_post_backward 完成梯度计算,无需手工编写反向逻辑。

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

【免费下载链接】ops-transformer

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

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

相关推荐

上一篇:Fiddler中文版性能优化:如何分析网站加载速度和瓶颈
下一篇:性能之巅:SyncTrayzor与主流Syncthing管理工具深度对比测试

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

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

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

立即咨询