CANN ops-transformer 算子详解:aclnnMhcPreSinkhorn 接口与 MhcPreSinkhorn 实现剖析
CANN ops-transformer 算子详解aclnnMhcPreSinkhorn 接口与 MhcPreSinkhorn 实现剖析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 开源仓库中 MHCMulti-Head Cross-attention / Hidden Compensation架构前向核心算子MhcPreSinkhorn展开完整讲解其aclnnMhcPreSinkhorn两段式 API 的功能语义、Sinkhorn 迭代归一化数学原理、全部入参/出参规格、平台约束与错误码并结合仓库源码给出 C 与 PyTorch 双路径的可运行调用示例。读完本文你将能够独立完成该算子在 Ascend NPU 上的单算子调用、理解其在训练反向传播中的中间量保存机制并掌握依据 shape/数据类型约束排查调用错误的技巧。1. 算子定位MHC 架构中的 Sinkhorn 归一化前向算子MhcPreSinkhorn 是mhc模块族中的核心前向算子。该模块族位于仓库 mhc 目录下包含mhc_pre、mhc_pre_sinkhorn、mhc_pre_sinkhorn_backward、mhc_sinkhorn、mhc_post等算子子目录共同支撑 MHC 网络的 pre/post 投影与 Sinkhorn 变换链路。本算子实现的功能可概括为对输入x执行 RmsNorm 归一化得到 $\vec{x}_{l}$计算三条投影路径pre / post / res的变换矩阵 $H^{pre}_l$、$H^{post}_l$、$H^{res}_l$其中 pre 路径经过 sigmoid、post 路径经过 $2\sigma$ 缩放res 路径进入 Sinkhorn 迭代对 $H^{res}_l$ 执行numIters轮 Sinkhorn 迭代归一化得到双随机矩阵doubly stochastic matrix作为最终 $H^{res}$输出 $h_{in}$Attention/MLP 层输入以及一组反向计算所需的中间结果。当needBackwardtrue时算子额外输出 sigmoid 之后的 $H^{pre}l$、$\vec{x}{l}$ 与 $\varphi$ 矩阵乘的结果、RmsNorm 的倒数 $\vec{x}_{l}$、迭代过程中的中间归一化结果normOut与求和结果sumOut这些中间张量将作为 MhcPreSinkhornBackward 反向算子的输入避免前向重复计算。2. 产品支持情况aclnnMhcPreSinkhorn接口及算子内核的产品支持矩阵如下产品是否支持Ascend 950PR / Ascend 950DT✅ 支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品✅ 支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品✅ 支持Atlas 200I/500 A2 推理产品❌ 不支持Atlas 推理系列产品❌ 不支持Atlas 训练系列产品❌ 不支持在源码层面算子注册于 mhc_pre_sinkhorn_def.cpp为ascend910b、ascend910_93添加 AICore 配置对应 Atlas A2/A3 系列为ascend950添加了独立配置并开启DynamicCompileStaticFlag、DynamicRankSupportFlag、DynamicShapeSupportFlag三项动态编译/动态 shape 能力扩展配置指向mhc_pre_sinkhorn_apt。由此可见Ascend 950 平台拥有独立的动态实现路径对应op_kernel下的mhc_pre_sinkhorn_apt.cpp与arch35内核支持能力也最完整见第 6 节规格差异。3. 功能与计算公式3.1 投影与归一化主公式设 $d$ 为尾轴隐藏维度算子按以下公式计算归一化向量与三条投影路径$$ \begin{aligned} \vec{x^{}{l}} \frac{1}{\sqrt{\frac{1}{d} \sum{\dim-2,\text{keepdim}\text{True}} x_i^2 \epsilon_{norm}}}\ H^{pre}l \alpha^{pre}{l} \cdot(\vec{x^{}{l}}\varphi^{pre}{l}) b^{pre}{l}\ H^{post}l \alpha^{post}{l} \cdot(\vec{x^{}{l}}\varphi^{post}{l}) b^{post}{l}\ H^{res}l \alpha^{res}{l} \cdot(\vec{x^{}{l}}\varphi^{res}{l}) b^{res}{l}\ H^{pre}l \sigma (H^{pre}{l})\ H^{post}l 2\sigma (H^{post}{l})\ h{in} \vec{x_{l}}H^{pre}_l \end{aligned} $$其中参数矩阵 $\varphi$ 沿行方向被拆分为三段使用$\varphi^{pre}$shape 为 $(N, N \times C)$、$\varphi^{post}$shape 为 $(N, N \times C)$、$\varphi^{res}$shape 为 $(N \times N, N \times C)$拼接后即 phi 的整体 shape $(N^2 2N, N \times C)$alpha 的 3 个元素依次对应 $\alpha^{pre}, \alpha^{post}, \alpha^{res}$bias 的 $2N N^2$ 个元素依次对应 $\beta^{pre}, \beta^{post}, \beta^{res}$。以上拆分约定可在 torchapi_mhc_pre_sinkhorn.md 的参数说明中找到原文佐证。3.2 Sinkhorn 迭代归一化以 $H^{res}_l$ 作为输入Sinkhorn 变换共执行numIters次迭代。迭代过程中交替对最后一维列方向与倒数第二维行方向做 softmax 式归一化生成中间结果normOut[k]与sumOut[k]最终输出最后一次迭代的normOut作为变换结果。第一次迭代初始化$$ \begin{aligned} \mathbf{normOut}[0] \text{softmax}(\mathbf{H^{res}l}, \dim-1) \epsilon{hc}, \ \mathbf{sumOut}[1] \sum_{\dim-2,\text{keepdim}\text{True}} \mathbf{normOut}[0] \epsilon_{hc}, \ \mathbf{normOut}[1] \frac{\mathbf{normOut}[0]}{\mathbf{sumOut}[1]}, \end{aligned} $$第 $i$ 次迭代$i 1, 2, \dots, \text{numIters}-1$$$ \begin{aligned} \mathbf{sumOut}[2i] \sum_{\dim-1,\text{keepdim}\text{True}} \mathbf{normOut}[2i-1] \epsilon_{hc}, \ \mathbf{normOut}[2i] \frac{\mathbf{normOut}[2i-1]}{\mathbf{sumOut}[2i]}, \ \mathbf{sumOut}[2i1] \sum_{\dim-2,\text{keepdim}\text{True}} \mathbf{normOut}[2i] \epsilon_{hc}, \ \mathbf{normOut}[2i1] \frac{\mathbf{normOut}[2i]}{\mathbf{sumOut}[2i1]}, \end{aligned} $$最终输出$$ \mathbf{normOut}[2 \times \text{numIters} - 1], \qquad \mathbf{sumOut}[2 \times \text{numIters} - 1] $$符号约定汇总符号含义$\mathbf{x}$输入张量MHC 层的 $H_{\text{res}}$ 矩阵$\epsilon_{hc}$ /hcEpsSinkhorn 迭代中的防除零参数$\epsilon_{norm}$ /normEpsRmsNorm 的防除零参数$\text{softmax}(\cdot, \dim-1)$在最后一维执行 softmax 归一化$\sum_{\dimd,\text{keepdim}\text{True}}$在指定维度 $d$ 上求和并保持维度$\mathbf{normOut}[k]$第 $k$ 步归一化中间结果$\mathbf{sumOut}[k]$第 $k$ 步求和中间结果$\mathbf{numIters}$迭代次数入参值得说明的是仓库中的 golden 实现 golden.py 对上述公式给出了可执行化的精确还原它对residual_logits先取行方向最大值做数值稳定的 softmaxrow_exp np.exp(residual_logits - row_max)随后按「行和→归一化→列和→归一化」的顺序交替迭代num_iters轮与文档公式一一对应该文件同时给出了hin np.sum(x_f32 * h_pre[..., :, None], axis-2)的 $h_{in}$ 计算方式可直接作为理解算子语义的参考基准。4. 两段式接口与函数原型每个算子遵循两段式接口规范必须先调用aclnnMhcPreSinkhornGetWorkspaceSize完成参数校验、构图并获取 workspace 大小与执行器再调用aclnnMhcPreSinkhorn执行实际计算。aclnnStatus aclnnMhcPreSinkhornGetWorkspaceSize( const aclTensor *x, const aclTensor *phi, const aclTensor *alpha, const aclTensor *bias, int64_t hcMult, int64_t numIters, double hcEps, double normEps, bool needBackward, aclTensor *hin, aclTensor *hPost, aclTensor *hRes, aclTensor *hPre, aclTensor *hcBeforeNorm, aclTensor *invRms, aclTensor *sumOut, aclTensor *normOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnMhcPreSinkhorn( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)从源码实现 aclnn_mhc_pre_sinkhorn.cpp 可以看到第一段接口内部的真实流程CheckParams依次执行空指针检查CheckNotNull、数据类型检查CheckDtypeValid、shape 检查CheckShape与格式检查CheckFormat若输入 tensor 为空x-IsEmpty()等直接返回成功且workspaceSize0对应文档「支持空 Tensor」语义对x、phi、alpha、bias分别插入l0op::Contiguous转连续算子即接口层面自动处理非连续 Tensor文档中标记为「√」的列构建l0op::MhcPreSinkhorn算子节点经uniqueExecutor-GetWorkspaceSize()取得 workspace 大小释放执行器给调用方供第二段接口使用。第二段接口aclnnMhcPreSinkhorn则直接调用CommonOpExecutorRun(workspace, workspaceSize, executor, stream)完成异步下发执行。5. aclnnMhcPreSinkhornGetWorkspaceSize 参数说明5.1 完整参数表第一段接口共 19 个参数4 个输入 Tensor、5 个属性、8 个输出 Tensor、2 个返回参数规格如下参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续 TensorxaclTensor*输入待计算数据表示网络中 mHC 层的输入数据支持空 TensorFLOAT16、BFLOAT16ND(bs, seq_len, n, c) / (t, n, c)√phiaclTensor*输入mHC 的参数矩阵支持空 TensorFLOAT32ND(n*n 2*n, n*c)√alphaaclTensor*输入mHC 的缩放参数支持空 TensorFLOAT32ND(3)√biasaclTensor*输入mHC 的 bias 参数支持空 TensorFLOAT32ND(n*n 2*n)√hcMultint64_t输入残差流数量HC 维度大小当前仅支持 4----numItersint64_t输入Sinkhorn 算法迭代次数当前仅支持 20----hcEpsdouble输入$H_{pre}$ 的 sigmoid 后的 eps 参数建议值1e-6----normEpsdouble输入RmsNorm 的防除零参数建议值1e-6----needBackwardbool输入是否需要输出额外属性中间结果建议值为 true----hinaclTensor*输出输出的 h_in作为 Attention/MLP 层的输入-FLOAT16、BFLOAT16ND(bs, seq_len, c) / (t, c)×hPostaclTensor*输出输出的 mHC 的 h_post 变换矩阵-FLOAT32ND(bs, seq_len, n) / (t, n)×hResaclTensor*输出输出的 mHC 的 h_res 变换矩阵-FLOAT32ND(bs, seq_len, n*n) / (t, n*n)×hPreaclTensor*输出需要反向时输出做完 sigmoid 计算之后的 hPre 矩阵needBackward 为 false 时此输出无效FLOAT32ND(bs, seq_len, n) / (t, n)×hcBeforeNormaclTensor*输出需要反向时输出x 与 phi 矩阵乘的结果needBackward 为 false 时此输出无效FLOAT32ND(bs, seq_len, n*n 2*n) / (t, n*n 2*n)×invRmsaclTensor*输出需要反向时输出RmsNorm 计算得到的 1/rneedBackward 为 false 时此输出无效FLOAT32ND(bs, seq_len, 1) / (t, 1)×sumOutaclTensor*输出需要反向时输出每一次迭代的 colSum/rowSum 结果needBackward 为 false 时此输出无效FLOAT32ND(sk_iter_count * 2, bs, seq_len, n) / (sk_iter_count * 2, t, n)×normOutaclTensor*输出需要反向时输出每一次 colSum/rowSum 迭代后的 comb 结果needBackward 为 false 时此输出无效FLOAT32ND(sk_iter_count * 2, bs, seq_len, n, n) / (sk_iter_count * 2, t, n, n)×workspaceSizeuint64_t*输出返回需要在 Device 侧申请的 workspace 大小-----executoraclOpExecutor**输出返回 op 执行器包含了算子计算流程-----格式约定4 维输入为 BSND 格式3 维输入为 TND 格式。其中bs表示 Batch 大小seq_len表示单 Batch 序列长度t表示所有 Batch 序列长度的累加和n表示残差流数量c表示尾轴隐藏维度大小。当n4、numIters20时hcMix n² 2n 24sumOut的 0 轴长度为 40normOut的 0 轴长度同样为 40。5.2 参数解析的源码依据hcMult / numIters 默认值与取值范围在 aclnn_mhc_pre_sinkhorn.cpp 中定义了HCMULT_DEFAULT_VALUE 4与NUM_ITERS_DEFAULT_VALUE 20CheckShape会对这两个属性做严格等值校验hcMult ! 4或numIters ! 20即返回ACLNN_ERR_PARAM_INVALID同时校验x的n维必须等于hcMult。算子定义 mhc_pre_sinkhorn_def.cpp 中对应的属性默认值亦为hc_mult4、num_iters20、hc_eps1e-6f、norm_eps1e-6f、need_backwardtrue。shape 依赖关系CheckShape依据x的 shape 推导 phi(hcMix, n*c)、alpha(3,)、bias(hcMix,)、hin(bs, seq_len, c)、hPost(bs, seq_len, n)、hRes(bs, seq_len, n*n)等所有输出的期望 shape任何一项不匹配都会报ACLNN_ERR_PARAM_INVALID。needBackward 开关CheckDtypeValid、CheckShape、CheckFormat中均以if (needBackward)分支决定是否校验 hPre、hcBeforeNorm、invRms、sumOut、normOut 这 5 个反向中间输出infershape 实现 mhc_pre_sinkhorn_infershape.cpp 中当needBackwardfalse时将这 5 个输出 shape 置为空{EMPTY_DIM}当needBackwardtrue时按第 5.1 节 shape 展开——这正是「输出无效/空 Tensor」语义的底层来源。空 Tensor 支持文档中 x、phi、alpha、bias 标注「支持空 Tensor」对应aclnnMhcPreSinkhornGetWorkspaceSize中的IsEmpty()短路逻辑空输入时直接返回成功且 workspace 为 0。6. 返回值与错误码aclnnMhcPreSinkhornGetWorkspaceSize与aclnnMhcPreSinkhorn均返回aclnnStatus状态码具体枚举可参考 aclnn 返回码。第一段接口完成入参校验以下场景会报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或必选属性且是空指针ACLNN_ERR_PARAM_INVALID161002输入变量的数据类型、shape 维度和数据格式不在支持的范围内ACLNN_ERR_PARAM_INVALID161002numIters 不为 20ACLNN_ERR_PARAM_INVALID161002n 值非 4ACLNN_ERR_RUNTIME_ERROR361001调用 NPU Runtime 接口申请内存/创建 Tensor 失败上述错误码的触发条件与 aclnn_mhc_pre_sinkhorn.cpp 中的CheckParams逻辑一一对应CheckNotNull失败返回ACLNN_ERR_PARAM_NULLPTRCheckDtypeValid/CheckShape/CheckFormat失败返回ACLNN_ERR_PARAM_INVALID其中hcMult与numIters的等值校验、n ! hcMult校验均可直接定位到CheckShape函数体。此外CheckShape还包含一条平台相关校验3 维 TND 输入仅 Ascend 950 平台支持在其他平台传入 3 维 x 会直接返回ACLNN_ERR_PARAM_INVALID。7. aclnnMhcPreSinkhorn 第二段接口参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口 aclnnMhcPreSinkhornGetWorkspaceSize 获取executor输入op 执行器包含了算子计算流程stream输入指定执行任务的 Stream返回值同样为aclnnStatus参见 aclnn 返回码。使用上要求workspace 需由用户通过aclrtMalloc在 Device 侧申请workspaceSize为 0 时可不申请执行前建议通过aclrtSynchronizeStream同步 Stream 以确保算子完成。8. 约束说明8.1 确定性计算与 Batch 一致性确定性计算aclnnMhcPreSinkhorn默认采用确定性实现相同输入多次调用结果一致。Batch 一致性Atlas A2 训练/推理系列、Atlas A3 训练/推理系列默认非 Batch 一致性实现不支持通过aclrtSetSysParamOpt开启 Batch 一致性Ascend 950PR/Ascend 950DT默认非 Batch 一致性实现支持通过aclrtSetSysParamOpt开启 Batch 一致性但开启后性能可能会有一定程度劣化。Batch 一致性的详细背景可参见 batch_consistency.md确定性计算语义可参见 determinism_compute.md。8.2 规格约束规格项规格规格说明numIters20目前只支持 20n4目前只支持 4Atlas A2 训练/推理系列、Atlas A3 训练/推理系列输入 x 仅支持 4 维 BSND 格式shape 为(bs, seq_len, n, c)尾轴 c 需为 128 的倍数且取值范围为[1, 100000]。Ascend 950PR / Ascend 950DT输入 x 支持 4 维 BSND 格式和 3 维 TND 格式shape 分别为(bs, seq_len, n, c)和(t, n, c)TND 格式下t 表示所有 Batch 序列长度的累加和对应输出 shape 按参数说明中的 t 轴展开尾轴 c 仅支持 4096、7168输入 phi 的数据范围限定在 $\pm \frac{1}{\sqrt{nc}}$ 范围内eps 限定在 1e-6此范围内具有较好的数值稳定性。最后一条约束在 golden 测试中也有体现golden.py 的_set_stable_inputs使用phi 1.0 / sqrt(hc_mult * d)填充 phi、alpha 0.25、bias 0.0构造数值稳定的测试输入与「phi 限定在 ±1/√(nc) 内、eps 取 1e-6」的规格约束相互印证。8.3 Tiling 与内核实现视角从算子 Host 侧 Tiling 实现 mhc_pre_sinkhorn_tiling.h 可以看到该算子在 A2/A3910B/910_93与 950 平台分别走两套 Tiling 数据布局MhcPreSinkhornTilingData含 BS 切分阈值BS_SPLIT_THRESHOLD、多核 M/K 切分、两阶段核心数等字段与MhcPreSinkhornRegbaseTilingData950 的 regbase 布局含 Sinkhorn 核数/行因子等字段。内核入口 mhc_pre_sinkhorn.cpp 展示了混合 AICAIV 的执行框架KERNEL_TYPE_MIX_AIC_1_2Tiling Key 0切 M阶段 1 的MhcPreSinkhornStage1完成 x 与 phi 的矩阵乘mm1_并产出 RmsNorm 的invRms与hcBeforeNorm阶段 2 的MhcPreSinkhornStage2完成 alpha/bias 融合、sigmoid、Sinkhorn 迭代与 h_in 输出Tiling Key 1切 MK按bsLoop分批循环每批由MhcPreSinkhornMembaseKSplitCorePart1/Part2两阶段处理批间通过SyncAllfalse()同步所有 AIV 核。这解释了为何算子能同时支持 BSND 与 TND 两种布局、以及在 BS 维上的切分调度能力也印证了needBackward中间输出sumOut/normOut等需要在 forward 阶段整体落盘以供反向算子使用的设计动机。9. 调用示例C / aclnn 两段式官方调用示例位于 examples/test_aclnn_mhc_pre_sinkhorn.cpp完整编译与执行流程请参考编译与运行样例。核心流程如下示例基于bs2, seq_len128, n4, c4096hcMix n² 2n 24num_iters20hc_epsnorm_eps1e-6need_backwardtrue#include iostream #include vector #include acl/acl.h #include aclnnop/aclnn_mhc_pre_sinkhorn.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) int64_t GetShapeSize(const std::vectorint64_t shape) { int64_t size 1; for (int64_t dim : shape) { size * dim; } return size; } // 创建输入 Tensorx 使用 BF16host 数据按 host_data[i] * 32767 量化到 int16 后转 uint16 存储 // phi/alpha/bias 使用 FLOAT32strides 按行主序计算格式均为 ACL_FORMAT_ND。 int CreateAclTensorBfloat16(const std::vectorfloat host_data, const std::vectorint64_t shape, void *device_addr, aclTensor *tensor) { int64_t size GetShapeSize(shape); std::vectoruint16_t host_data_bf16(size); for (int64_t i 0; i size; i) { host_data_bf16[i] static_castuint16_t(static_castint16_t(host_data[i] * 32767)); } int64_t byte_size size * sizeof(uint16_t); aclError ret aclrtMalloc(device_addr, byte_size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed, error: %d\n, ret); return -1); ret aclrtMemcpy(device_addr, byte_size, host_data_bf16.data(), byte_size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMemcpy failed, error: %d\n, ret); return -1); std::vectorint64_t strides(shape.size(), 1); for (int64_t i shape.size() - 2; i 0; --i) { strides[i] strides[i 1] * shape[i 1]; } tensor aclCreateTensor(shape.data(), shape.size(), ACL_BF16, strides.data(), 0, ACL_FORMAT_ND, shape.data(), shape.size(), device_addr); CHECK_RET(tensor ! nullptr, LOG_PRINT(aclCreateTensor failed\n); return -1); return 0; } // 输出 Tensor 创建无需 strides直接按 shape 创建hin 为 BF16其余为 FLOAT32 int CreateAclTensorFloat32Output(const std::vectorint64_t shape, void *device_addr, aclTensor *tensor) { int64_t size GetShapeSize(shape); int64_t byte_size size * sizeof(float); aclError ret aclrtMalloc(device_addr, byte_size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed, error: %d\n, ret); return -1); tensor aclCreateTensor(shape.data(), shape.size(), ACL_FLOAT, nullptr, 0, ACL_FORMAT_ND, shape.data(), shape.size(), device_addr); CHECK_RET(tensor ! nullptr, LOG_PRINT(aclCreateTensor failed\n); return -1); return 0; } int main() { int32_t device_id 0; aclrtContext context nullptr; aclrtStream stream nullptr; int64_t bs 2, seq_len 128, n 4, c 4096; int64_t hc_mult n; int64_t hc_mix n * n 2 * n; int64_t num_iters 20; float hc_eps 1e-6f, norm_eps 1e-6f; bool need_backward true; std::vectorint64_t x_shape {bs, seq_len, n, c}; std::vectorint64_t phi_shape {hc_mix, n * c}; std::vectorint64_t alpha_shape {3}; std::vectorint64_t bias_shape {hc_mix}; std::vectorint64_t hin_shape {bs, seq_len, c}; std::vectorint64_t h_post_shape {bs, seq_len, n}; std::vectorint64_t h_res_shape {bs, seq_len, n * n}; std::vectorint64_t h_pre_shape {bs, seq_len, n}; std::vectorint64_t hc_before_norm_shape {bs, seq_len, hc_mix}; std::vectorint64_t inv_rms_shape {bs, seq_len, 1}; std::vectorint64_t sum_out_shape {num_iters * 2, bs, seq_len, n}; std::vectorint64_t norm_out_shape {num_iters * 2, bs, seq_len, n, n}; // 1. 初始化 ACLaclInit - aclrtSetDevice - aclrtCreateContext - // aclrtSetCurrentContext - aclrtCreateStream aclError ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed, error: %d\n, ret); return -1); ret aclrtSetDevice(device_id); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed, error: %d\n, ret); return -1); ret aclrtCreateContext(context, device_id); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateContext failed, error: %d\n, ret); return -1); ret aclrtSetCurrentContext(context); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetCurrentContext failed, error: %d\n, ret); return -1); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed, error: %d\n, ret); return -1); // 2. 创建输入/输出 Tensor此处省略各创建函数的展开逻辑同 CreateAclTensorBfloat16 / // CreateAclTensorFloat32Outputphi/alpha/bias 使用 FLOAT32 创建 void *x_addr, *phi_addr, *alpha_addr, *bias_addr; aclTensor *x, *phi, *alpha, *bias; // ... 按上述 shape 创建 4 个输入 Tensor ... void *hin_addr, *h_post_addr, *h_res_addr; aclTensor *hin, *h_post, *h_res; // ... 按上述 shape 创建 8 个输出 Tensorhin 为 BF16其余为 FLOAT32... // 3. 第一段接口获取 workspace 大小与执行器 uint64_t workspace_size 0; aclOpExecutor *executor nullptr; aclnnStatus aclnn_ret aclnnMhcPreSinkhornGetWorkspaceSize( x, phi, alpha, bias, hc_mult, num_iters, hc_eps, norm_eps, need_backward, hin, h_post, h_res, h_pre, hc_before_norm, inv_rms, sum_out, norm_out, workspace_size, executor); CHECK_RET(aclnn_ret ACL_SUCCESS, LOG_PRINT(aclnnMhcPreSinkhornGetWorkspaceSize failed, error: %d\n, aclnn_ret); return -1); // 4. 申请 workspace 并执行第二段接口 void *workspace_addr nullptr; if (workspace_size 0) { ret aclrtMalloc(workspace_addr, workspace_size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc workspace failed, error: %d\n, ret); return -1); } aclnn_ret aclnnMhcPreSinkhorn(workspace_addr, workspace_size, executor, stream); CHECK_RET(aclnn_ret ACL_SUCCESS, LOG_PRINT(aclnnMhcPreSinkhorn failed, error: %d\n, aclnn_ret); return -1); // 5. 同步 Stream 确保计算完成随后销毁 Tensor、释放 Device 内存并 aclFinalize CHECK_RET(aclrtSynchronizeStream(stream) ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed\n); return -1); LOG_PRINT(MhcPreSinkhorn compute success!\n); // 资源释放aclDestroyTensor - aclrtFree - aclrtDestroyStream - aclrtDestroyContext // - aclrtResetDevice - aclFinalize return 0; }运行示例的输出 shape 与第 5.1 节规格完全一致sum_out_shape的 0 轴为num_iters * 2 40norm_out_shape为(40, bs, seq_len, 4, 4)。示例完整代码含InitAcl、CreateInputTensors、CreateOutputTensors、PrintTensorDataFloat/PrintTensorDataBfloat16调试打印与全部资源释放逻辑可直接参考仓库 test_aclnn_mhc_pre_sinkhorn.cpp。10. PyTorch 调用路径与反向联动除 aclnn 接口外该算子还提供 PyTorch 封装其 API 文档见 torchapi_mhc_pre_sinkhorn.mdPython 入口实现于 torch_extension/mhc_pre_sinkhorn.py。10.1 单算子模式调用import torch import torch_npu from cann_ops_transformer.ops import mhc_pre_sinkhorn B 1 S 1024 N 4 C 4096 x torch.randn(B, S, N, C, dtypetorch.bfloat16).npu() phi torch.randn(N * N 2 * N, N * C, dtypetorch.float32).npu() alpha torch.randn(3, dtypetorch.float32).npu() bias torch.randn(N * N 2 * N, dtypetorch.float32).npu() hcMult 4 numIters 20 hcEps 1e-6 normEps 1e-6 hin, hPost, hRes mhc_pre_sinkhorn( x, phi, alpha, bias, hcMult, numIters, hcEps, normEps )10.2 图模式调用import torch import torch_npu import torchair from cann_ops_transformer.ops import mhc_pre_sinkhorn torch_npu.npu.set_device(0) B 1 S 128 N 4 C 4096 class MhcPreSinkhornModel(torch.nn.Module): def forward(self, x, phi, alpha, bias): hin, hPost, hRes mhc_pre_sinkhorn( x, phi, alpha, bias, hc_mult4, num_iters20, hc_eps1e-6, norm_eps1e-6 ) return hin, hPost, hRes model MhcPreSinkhornModel().npu() npu_backend torchair.get_npu_backend() model torch.compile(model, backendnpu_backend, dynamicFalse) x torch.randn(B, S, N, C, dtypetorch.bfloat16, devicenpu) phi torch.randn(N * N 2 * N, N * C, dtypetorch.float32, devicenpu) alpha torch.randn(3, dtypetorch.float32, devicenpu) bias torch.randn(N * N 2 * N, dtypetorch.float32, devicenpu) hin, hPost, hRes model(x, phi, alpha, bias)PyTorch 路径同样支持 TND 格式3 维 x且 torch API 文档中的输出 shape 表hin/hPost/hRes与约束A3/A2 平台仅 BSND、尾轴 128 对齐950 平台支持 TND、尾轴仅 4096/7168与 aclnn 接口保持一致。10.3 autograd 中间量保存机制从 mhc_pre_sinkhorn.py 的实现可以看到它与底层needBackward标志的联动设计mhc_pre_sinkhorn(x, phi, alpha, bias, hc_mult, num_iters, hc_eps1e-6, norm_eps1e-6)对外只返回(hin, h_post, h_res)三个主输出若x/phi/alpha/bias任一设置了requires_gradTrue则走MhcPreSinkhornFunctiontorch.autograd.Function路径forward中以out_flagTrue调用底层算子得到全部 8 个输出并通过ctx.save_for_backward保存h_pre, hc_before_norm, inv_rms, sum_out, norm_out等中间量backward则调用 mhc_pre_sinkhorn_backward 模块完成对grad_hin, grad_h_post, grad_h_res的反向传播若无需梯度则走普通路径out_flagFalse底层不计算/不返回 5 个中间输出以节省内存——这与第 5.1 节中「needBackward 为 false 时中间输出无效」的语义完全对应。11. 验证与测试资源仓库为该算子提供了多维度的验证资源可作为二次开发与自测的参考Golden 参考实现tests/assets/golden.py 提供 NumPy/Torch 版mhc_pre_sinkhorn_goldenaclnn 与 kernel 两条入口均注册完整复现 RmsNorm、投影、sigmoid、Sinkhorn 迭代与h_in计算并内置数值稳定的输入生成器_set_stable_inputsST 用例配置tests/st/ttk_aclnn_mhc_pre_sinkhorn.csv 与 tests/st/ttk_kernel_mhc_pre_sinkhorn.csv 覆盖 aclnn 接口级与 kernel 级用例Tiling UTtests/ut/op_host/test_mhc_pre_sinkhorn_tiling.cpp 覆盖 Host 侧 Tiling 计算逻辑GEIR 示例examples/test_geir_mhc_pre_sinkhorn.cpp 提供图引擎GEIR调用路径示例。12. 小结aclnnMhcPreSinkhorn是 MHC 网络在 NPU 上高效落地 Sinkhorn 归一化的关键前向算子其核心价值在于将 RmsNorm、三条投影、sigmoid 缩放与 Sinkhorn 迭代融合为单一算子避免中间张量反复进出 Global Memory同时通过needBackward开关按需产出hPre/hcBeforeNorm/invRms/sumOut/normOut五类中间结果为反向算子提供零重复计算的梯度输入。使用时的关键要点可归纳为严格遵循两段式接口先GetWorkspaceSize后执行且 workspace 须在 Device 侧申请属性约束硬性固定hcMult必须为 4、numIters必须为 20shape 严格依赖n与cA3/A2 平台 c 需 128 对齐950 平台 c 仅 4096/7168平台差异明显TND 格式仅 Ascend 950 支持Batch 一致性仅 Ascend 950 可通过aclrtSetSysParamOpt开启训练场景建议needBackwardtruetorch 封装在需要梯度时自动启用推理场景可置 false 以节省内存在 Ascend 950 上使用 phi 时建议将数据控制在 $\pm 1/\sqrt{nc}$ 范围内并保持 eps1e-6以获得较好的数值稳定性。相关文档与源码入口接口文档 aclnnMhcPreSinkhorn.md、PyTorch 文档 torchapi_mhc_pre_sinkhorn.md、模块 README mhc_pre_sinkhorn/README.md、两段式接口规范 two_phase_api.md、aclnn 返回码 aclnn_return_code.md、编译运行样例 compile_and_run_sample.md。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →