CANN 自定义算子 npu_hc_pre_sinkhorn 深度解析:mHC 结构 Sinkhorn 归一化融合算子实战指南

发布时间:2026/9/18 16:44:30
CANN 自定义算子 npu_hc_pre_sinkhorn 深度解析:mHC 结构 Sinkhorn 归一化融合算子实战指南 CANN 自定义算子 npu_hc_pre_sinkhorn 深度解析mHC 结构 Sinkhorn 归一化融合算子实战指南【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer本指南以 CANN AscendC 自定义算子custom.npu_hc_pre_sinkhorn为对象完整讲解其在 mHCMulti-Head Combination注意力前处理流程中的作用、函数原型、全部参数/返回值约束、底层 tiling 与 kernel 实现原理并结合仓库内的 NumPy 参考实现与精度测试用例给出可直接运行的调用示例。读完本文你将掌握该算子的数学语义、Host/Device 两侧实现脉络以及如何基于 test_npu_hc_pre_sinkhorn.py 完成端到端验证。一、算子定位mHC 前处理链路中的 Sinkhorn 环节hc_pre_sinkhorn是 mHCMulti-Head Combination结构前处理部分的融合算子。在完整链路中mHC 前处理被拆分为两个可独立执行的 AscendC 小算子npu_hc_pre_inv_rms计算输入x的 InvRms1 / RMS(x)输出 shape 为[T, 1]或[b, s, 1]作为后续混合权重缩放的归一化系数npu_hc_pre_sinkhorn本文主角负责 hc_pre 中 Sinkhorn 部分的全部计算包括pre、post混合权重的 sigmoid 生成、comb_frag组合矩阵的 Sinkhorn 归一化以及基于加权组合对x的融合输出y。从仓库文档的约束说明可知该拆分方案是 mHC 融合算子的降级路径当T/bs大于 128 或不能被 16 整除时会启用hc_pre_inv_rmshc_pre_sinkhorn小算子拼接执行而在T/bs ≤ 128且能被 16 整除时则直接走高性能的 npu_hc_pre 融合算子详见 custom-npu_hc_pre.md 的约束说明。因此在 decode 场景的 bs 拆分与动态 shape 适配中hc_pre_sinkhorn承担着重要的兜底计算职责。产品支持情况产品是否支持Atlas A3 推理系列产品√Ascend 950PR/Ascend 950DT√从源码 hc_pre_sinkhorn_def.cpp 可印证算子级配置注册了ascend910b、ascend910_93对应 Atlas A3 系列以及ascend950Ascend 950 系列三档 AICore 配置其中ascend950采用独立的 regbase 配置开启动态编译、动态 rank 与动态 shape 支持。二、功能与计算流程从 mixes 到 y/post/comb_fraghc_pre_sinkhorn的计算输入为四路张量mixes混合系数 logits、rsqrtInvRms 结果、hc_scale缩放系数、hc_base偏移系数以及待加权组合的隐状态x。其整体计算语义与仓库测试用例中的 NumPy 参考实现 test_npu_hc_pre_sinkhorn.py 完全对应可分四步理解第 1 步rsqrt 缩放。将mixes按最后一维乘上rsqrt相当于对每行做 RMS 归一化的逆缩放得到参与后续计算的混合 logitsmixes mixes * rsqrt第 2 步pre 权重加权系数计算。取mixes的[0 : hc_mult]段配合hc_scale[0]与hc_base[0 : hc_mult]做仿射变换后过 sigmoid再加hc_eps防零pre sigmoid(mixes0 * hc_scale[0] hc_base0) hc_eps第 3 步post 权重计算。取mixes的[hc_mult : 2*hc_mult]段同样仿射 sigmoid再乘 2post sigmoid(mixes1 * hc_scale[1] hc_base1) * 2第 4 步comb_frag 组合矩阵与 Sinkhorn 迭代。将mixes的[2*hc_mult : ]段 reshape 为[..., hc_mult, hc_mult]仿射变换后先做带hc_eps的 softmax沿最后一维随后进入 Sinkhorn 行-列交替归一化迭代共hc_sinkhorn_iters - 1轮comb_frag softmax(mixes2 * hc_scale[2] hc_base2, -1) hc_eps comb_frag comb_frag / (sum(comb_frag, axis-2) hc_eps) # 列归一化 for _ in range(hc_sinkhorn_iters - 1): comb_frag comb_frag / (sum(comb_frag, axis-1) hc_eps) # 行归一化 comb_frag comb_frag / (sum(comb_frag, axis-2) hc_eps) # 列归一化第 5 步y 融合输出。将pre扩展一维作为权重乘到x上再沿 hc 轴hc_mult维求和得到最终的组合结果y sum(pre[..., None] * x, axis2)整体输出为三个张量y加权组合后的隐状态、post后混合权重、comb_fragSinkhorn 归一化后的组合矩阵。其中comb_frag与post会作为 mHC 后续阶段如 hc_post 融合算子的输入继续参与计算。三、函数原型custom.npu_hc_pre_sinkhorn(Tensor mixes, Tensor rsqrt, Tensor hc_scale, Tensor hc_base, Tensor x, int hc_mult4, int hc_sinkhorn_iters20, float hc_eps1e-5) - (Tensor, Tensor, Tensor)算子注册名在 Torch 侧的调用方式为torch.ops.custom.npu_hc_pre_sinkhorn(...)与源码 hc_pre_sinkhorn_def.cpp 中OP_ADD(HcPreSinkhorn)注册的HcPreSinkhorn算子对应。四、参数说明说明bbatch size表示输入样本批量大小、ssequence length表示输入样本序列长度、hchead count表示注意力头数、dhead dimension表示注意力头的维度数、T 表示 bs 合轴后的大小。mixesTensor必选参数输入 tensor。不支持非连续数据格式支持 ND数据类型支持floatshape 为[T, hc_mix]或[b, s, hc_mix]。混合系数 logits经 rsqrt 缩放后按段切分生成 pre/post/comb_frag 三类权重。rsqrtTensor必选参数输入 tensor。不支持非连续数据格式支持 ND数据类型支持floatshape 为[T, 1]或[b, s, 1]。即hc_pre_inv_rms算子的输出用于对 mixes 做 RMS 逆归一化缩放。hc_scaleTensor必选参数输入 tensor。不支持非连续数据格式支持 ND数据类型支持floatshape 为[3]。三段混合权重各自使用的缩放系数tiling 侧会强校验其第一维必须等于 3见 hc_pre_sinkhorn_tiling.cpp。hc_baseTensor必选参数输入 tensor。不支持非连续数据格式支持 ND数据类型支持floatshape 为[hc_mix]。偏移系数tiling 侧强校验其第一维必须等于 hc_mix见 hc_pre_sinkhorn_tiling.cpp。xTensor必选参数输入 tensor。不支持非连续数据格式支持 ND数据类型支持bfloat16shape 为[T, hc_mult, d]或[b, s, hc_mult, d]。mHC 结构待组合的隐状态。hc_multint固定为 4表示 hc 头数合并的倍数。hc_sinkhorn_itersint可选Sinkhorn 迭代次数取值固定为 20。hc_epsfloat可选计算过程中的 ε 参数Host 侧参数仅支持 double 类型默认值为 1e-05。用于 sigmoid 输出与 softmax/归一化分母的数值稳定处理。五、返回值说明yTensor输出 tensor。数据格式支持 ND数据类型支持bfloat16shape 为[T, d]或[b, s, d]。pre 加权组合后的融合隐状态。postTensor输出 tensor。数据格式支持 ND数据类型支持floatshape 为[T, hc_mult]或[b, s, hc_mult]。sigmoid 后乘 2 的后混合权重。comb_fragTensor输出 tensor。数据格式支持 ND数据类型支持floatshape 为[T, hc_mult, hc_mult]或[b, s, hc_mult, hc_mult]。经 Sinkhorn 迭代归一化的组合矩阵。三个输出的数据类型与算子定义文件 hc_pre_sinkhorn_def.cpp 完全一致y为DT_BF16post与comb_frag为DT_FLOAT。六、约束说明shape 字段取值范围约束字段名取值规则与说明hc_mult取值固定为: 4d取值固定为4096hc_mix取值固定为: 24该接口支持推理场景下使用。该接口支持 aclgraph 入图。该接口与 PyTorch 配合使用时需要保证 CANN 相关包与 PyTorch 相关包的版本匹配。输入 tensor 均要求连续内存不支持非连续输入shape 同时支持 2 维合轴形态[T, ...]与 3 维 b/s 形态[b, s, ...]两种表达tiling 侧会依据mixes的维度数自动区分并做bs b * s的合轴换算见 hc_pre_sinkhorn_tiling.cpp。七、源码级实现Host 侧 tiling 与 Device 侧 kernel7.1 算子定义OpDefhc_pre_sinkhorn_def.cpp 中5 个输入mixes/rsqrt/hc_scale/hc_baseDT_FLOAT与xDT_BF16全部为 REQUIRED格式统一为FORMAT_ND三个属性hc_mult默认 4、hc_sinkhorn_iters默认 20、hc_eps默认 1e-6f均为 OPTIONAL。kernel 文件按编译宏切分实现__DAV_C310__Ascend 950 系列走 regbase 版hc_pre_sinkhorn_regbase_perf.h/hc_pre_sinkhorn_regbase_base.h其余平台走 membase 版hc_pre_sinkhorn_perf.h/hc_pre_sinkhorn_base.h见 hc_pre_sinkhorn.cpp。7.2 Host 侧 Tiling 策略tiling 实现位于 hc_pre_sinkhorn_tiling.cpp核心要点多核行切分按bs在 AIV core 之间均分rowOfFormerBlock_ CeilDiv(bs, coreNum)并计算首块与尾块行数rowOfFormerBlock_/rowOfTailBlock_用于处理不整除场景。UB 容量驱动的二维分块对rowFactor_单次搬入的 bs 行数与dFactor_单次处理的 d 分片做联合规划。若 mix/rsqrt/post/comb_frag 等小张量加上x/y能整体放入 UB则dLoop_ 1全量加载否则按 2 的幂逐步减小dFactor_直至x、y的 double-buffer 占用满足 UB 剩余空间dFactor_ 32时按 32 向下对齐。若d可全载则反向增大rowFactor_以搬运更多行。平台差异ASCEND950regbase 架构与其余平台分别走CalcRegbaseOpTiling/CalcMembaseOpTiling二者在xCastSize、yCastSize、广播 buffer 与 reduce buffer 等临时空间的预算上有所区别。输出tiling data 记录bs、hcMix、hcMult、d、各分块因子与循环次数、iterTimes、eps并设置SetBlockDim(usedCoreNums_)与 workspace 大小32 字节最终写入 raw tiling data 供 kernel 读取。7.3 Device 侧 Kernel 核心算子kernel 公共实现位于 hc_pre_sinkhorn_base.h关键计算单元与数学流程一一对应逐行广播二元运算模板MulABLastDimBrcInline/SubABLastDimBrcInline/DivABLastDimBrcInline通过Brcb将末维向量广播为行矩阵再按列宽与行数决策在 Col 方向或 Row 方向开 RepeatRepeat 数上限 255同时处理numRepeatPerLine与numRemainPerLine的尾块AddBAFirstDimBrcInline处理首维行向广播加法。Sigmoid 实现SigmoidPerf通过Muls(-1)→Exp→Adds(1)计算分母再以 broadcast 除法得到1 / (1 exp(-x))对应公式中的 sigmoid。pre/post 处理ProcessPre完成mixes * rsqrt→ 乘scale→ 加hc_base→ sigmoid → 加eps的全流程ProcessPost与之类似但在 sigmoid 后额外Muls(2.0f)。Softmax 与 Sinkhorn 迭代SoftmaxFP32Perf使用WholeReduceMax求行内最大值做数值稳定减法再Exp、WholeReduceSum、Div最后加epsDivABABrcInline实现[bs, hc_mult, hc_mult]对[bs, 1, hc_mult]的广播除法用于行/列归一化交替迭代。注释明确当前支持 R 轴即hc_mult维小于等于 64 的场景与固定hc_mult4的约束吻合。y 融合ProcessY将 bf16 的xCast 到 floatCastTwoDim与pre权重做逐行广播乘再经ReduceSumARAPerf沿 hc 轴dim1归约求和最后 Cast 回 bf16 输出。数据搬运CopyIn/CopyOut基于DataCopyPad/DataCopyPadExtParams实现CopyInWithOuterFor以 4 维张量视角逐外层循环搬运comb_frag等输入。从 kernel 结构可以看出整条计算链路在 Device 侧以 float32 为主计算精度仅在输入输出边界做 bf16 转换与测试用例中x 转 float32 计算、y 回 bf16的参考实现策略一致。八、调用示例与精度验证8.1 端到端调用示例仓库提供了可直接运行的完整示例 test_npu_hc_pre_sinkhorn.py。其主流程为构造随机输入 → 用 NumPy 参考实现计算 CPU 基准 → 将输入搬到 NPU 调用torch.ops.custom.npu_hc_pre_sinkhorn→ 再经torch.compiletorchair后端reduce-overhead模式走图模式执行 → 三者逐项对比精度。核心调用代码import torch import torch_npu import torchair import custom_ops DEVICE_ID 0 torch_npu.npu.set_device(int(DEVICE_ID)) # 参数 b, s, hc_mix, hc_mult 4, 4, 24, 4 d 4096 # 示例还覆盖了 d 7168 的对比场景 hc_sinkhorn_iters, hc_eps 20, 1e-6 # 构造输入与文档约束一致mixes/rsqrt/hc_scale/hc_base 为 float32x 为 bfloat16 mixes torch.tensor(np.random.uniform(-2, 2, (b, s, hc_mix))).to(torch.float32) rsqrt torch.tensor(np.random.uniform(-2, 2, (b, s))).to(torch.float32) hc_scale torch.tensor(np.random.uniform(-2, 2, (3))).to(torch.float32) hc_base torch.tensor(np.random.uniform(-2, 2, (hc_mix))).to(torch.float32) x torch.tensor(np.random.uniform(-2, 2, (b, s, hc_mult, d))).to(torch.bfloat16) # NPU 直调Eager 模式 npu_yOut, npu_postOut, npu_comb_fragOut torch.ops.custom.npu_hc_pre_sinkhorn( mixes.to(npu:0), rsqrt.to(npu:0), hc_scale.to(npu:0), hc_base.to(npu:0), x.to(npu:0), hc_mult, hc_sinkhorn_iters, hc_eps) # torch.compile 图模式aclgraph 入图路径 from torchair.configs.compiler_config import CompilerConfig config CompilerConfig() config.mode reduce-overhead npu_backend torchair.get_npu_backend(compiler_configconfig) npu_mode torch.compile(Network().to(npu:0), fullgraphTrue, backendnpu_backend, dynamicFalse) npu_yOut, npu_postOut, npu_comb_fragOut npu_mode(mixes_npu, rsqrt_npu, hc_scale_npu, hc_base_npu, x_npu, hc_mult, hc_sinkhorn_iters, hc_eps)示例中Network仅是一个将算子调用包装在forward中的nn.Module便于torch.compile捕获成图Eager 与图模式两种调用路径的输入、输出语义完全一致。8.2 精度验证口径测试用例对三路输出分别做精度断言见 test_npu_hc_pre_sinkhorn.pyy与 CPU 参考对比rtolatol1e-4bf16 输出容差较宽post与comb_fragrtolatol1e-5float 输出容差更严。同时测试覆盖了d ∈ {4096, 7168}两种维度且 CPU 参考计算全程使用float64以保证基准精度。需要说明的是示例中的hc_eps取 1e-6而算子属性的默认值为 1e-5两者均可显式传入hc_eps仅支持 double 类型Python float 即满足。九、与相邻算子的协同关系hc_pre_sinkhorn不是孤立算子它在 mHC 前处理链路中与其他算子形成完整闭环上游npu_hc_pre_inv_rms 输出的rsqrt张量直接作为本算子的第二输入InvRms 公式为InvRms(x) 1 / RMS(x)其中RMS(x) sqrt(mean(x²) eps)。下游本算子输出的post、comb_frag作为 mHC 后处理部分的输入文档目录中的 custom-npu_hc_post.md 记录了对应的后处理融合算子。整体npu_hc_pre 提供npu_hc_pre/npu_hc_pre_v2两个融合入口内部同时包含 InvRms 与 Sinkhorn 计算当 shape 条件不满足融合使能条件T/bs 128或不可被 16 整除时才退化为hc_pre_inv_rmshc_pre_sinkhorn的拼接路径。因此在动态 shape、超长序列等场景下理解本算子的行为对保障 mHC 推理正确性与性能至关重要。十、使用建议与注意事项shape 合法性d固定为 4096、hc_mix固定为 24、hc_mult固定为 4hc_scale长度必须为 3hc_base长度必须与hc_mix一致Host 侧 tiling 会直接报错拦截非法输入。内存连续性所有输入均不支持非连续张量调用前需确保张量内存连续避免transpose、切片等产生非连续视图后直接传入。版本配套与 PyTorch 配合使用如示例中的torch.compile torchair 后端时需保证 CANN 相关包与 PyTorch 相关包的版本匹配并按需安装torch_npu、torchair、custom_ops依赖。精度选择hc_eps只接受 double 类型注意以 Python float 传入数值上建议与上下游hc_pre 融合算子默认 1e-6保持一致避免归一化行为不一致。性能路径认知本算子属于小算子拼接的降级路径若你的推理 shape 满足T/bs ≤ 128且能被 16 整除优先评估是否可直接使能npu_hc_pre融合算子以获得更优性能。十一、总结npu_hc_pre_sinkhorn是 CANN 生态中面向 mHC 注意力结构的高价值 AscendC 融合算子将 sigmoid 权重生成、softmax 与 Sinkhorn 迭代归一化、加权组合三个计算环节融合进单个算子并以 float32 内核算精度、bf16 边界 IO 的方式兼顾精度与带宽。本文从函数原型、参数约束、数学语义、Host/Device 双层实现到端到端精度验证做了全链路解析读者可结合 算子文档、测试示例 与 源码目录 进一步深入定制与调优。【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考