资讯动态

CANN ops-transformer 融合门控 Delta 网络解码算子 FusedGdnDecode:功能原理与 aclnn/torch 双接口实战指南

发布时间:2026/9/19 22:41:17 来源:尧图企业网站定制
CANN ops-transformer 融合门控 Delta 网络解码算子 FusedGdnDecode功能原理与 aclnn/torch 双接口实战指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerFusedGdnDecode 是 CANN ops-transformer 算子库中面向门控 Delta 网络Gated Delta NetworkGDN推理解码场景的高性能融合算子将单 token 解码过程中的 QKV 拆分、Q/K 归一化、门控计算、循环状态更新与输出投影合并为一次 NPU 执行。本文基于 experimental/attention/fused_gdn_decode/README.md 与其配套的 aclnn 接口文档、调用示例、torch 扩展及 host/kernel 源码系统讲解该算子的产品支持范围、数学语义、输入输出约束、两段式 aclnn 调用方法、torch 接口使用方式并深入剖析其 tiling 与 kernel 的底层实现帮助读者在 Atlas A2/A3 系列产品上快速完成集成与调优。产品支持情况FusedGdnDecode 算子对当前 CANN 产品的支持矩阵如下数据来源README.md产品是否支持Ascend 950PR/Ascend 950DT×Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×在算子定义注册文件 op_host/fused_gdn_decode_def.cpp 中OpAICoreConfig通过AddConfig(ascend910b, ...)与AddConfig(ascend910_93, ...)绑定 AICore 配置从源码侧印证了该算子面向 Ascend 910B 与 Ascend 910_93Atlas A2/A3 系列平台适配的定位。功能说明算子职责FusedGdnDecode 完成门控 Delta 网络单 token 解码计算将以下五个步骤融合为一次算子执行QKV 拆分从拼接张量 mixedQkv 中按 Q、K、V 顺序拆分出对应的切片Q/K 归一化对拆分出的 q 与 k 分别做 RMS 归一化门控计算根据门控输入 a、b、aLog、dtBias 计算衰减门控 g 与输出门控 β循环状态更新按 GDN 的状态递推公式原地更新循环状态矩阵 S输出投影利用更新后的状态计算当前 token 的输出 out。计算公式Q/K 归一化ε 为数值稳定小量kernel 源码中取EPS 1.0e-6f见 op_kernel/fused_gdn_decode.h$$ q_h scale \times \frac{q_h}{\sqrt{\sum_i q_{h,i}^2\epsilon}}, \quad k_h \frac{k_h}{\sqrt{\sum_i k_{h,i}^2\epsilon}} $$门控计算$$ g_j -\exp(A_j)\times softplus(a_jdtBias_j), \quad \beta_j sigmoid(b_j) $$循环状态更新与输出$$ S_j \exp(g_j)S_j\beta_j(v_j-\exp(g_j)S_jk_h)k_h^T, \quad out_j S_jq_h $$其中$h\lfloor j/(H_v/H)\rfloor$$S_j\in R^{V\times K}$stateRef原地保存更新后的状态。从数学形式可以直观看出算子的核心计算模式exp(g_j)S_jk_h是一个矩阵-向量乘MatVecMul(v_j - ·)k_h^T构成一次秩一更新Rank-One UpdateS_jq_h再次是矩阵-向量乘最后通过按行归约ReduceSum得到输出。kernel 源码 op_kernel/fused_gdn_decode.h 中的MatVecMul、RankOneUpdate、ReduceRows三个内联函数正是对这些计算的向量化实现。参数说明参数名输入/输出/属性描述数据类型数据格式mixedQkv输入按 Q、K、V 顺序拼接的输入shape 为(B, 2*H*KHv*V)。BFLOAT16、FLOAT16NDa输入门控输入 ashape 为(B, Hv)。与 mixedQkv 一致NDb输入门控输入 bshape 为(B, Hv)。与 mixedQkv 一致NDaLog输入门控参数 Ashape 为(Hv,)。FLOAT32NDdtBias输入softplus 偏置shape 为(Hv,)。与 mixedQkv 一致NDstateRef输入输出循环状态shape 为(BlockNum, Hv, V, K)原地更新。FLOAT32 或与 mixedQkv 一致NDssmStateIndices输入batch 到 stateRef 槽位的映射shape 为(B,)。INT32NDscale属性Q 归一化后的缩放系数默认值为 1.0。FLOAT-softplusThreshold属性softplus 阈值默认值为 20.0。FLOAT-out输出GDN 输出shape 为(B, 1, Hv, V)。与 mixedQkv 一致ND参数语义补充说明H 与 HvmixedQkv 的拼接维度2*H*K表明 Q、K 各含 H 个头、每个头维度为 K而门控与状态均按 Hv 个头组织。tiling 源码 op_host/fused_gdn_decode_tiling.cpp 会从 mixedQkv 最后一维反解出 Hh qkDim / (2*k)并要求hv % h 0每个 Q/K 头被hv/h个 V 头共享对应公式中的 $h\lfloor j/(H_v/H)\rfloor$。stateRef 的原地更新语义状态按BlockNum个槽位组织ssmStateIndices[i]决定第 i 个 batch 使用哪个槽位更新后的状态写回原槽位供下一 token 解码继续使用是 GDN 这类循环线性注意力模型实现流式解码的关键。scale 与 softplusThreshold 属性默认值两个属性在 op_host/fused_gdn_decode_def.cpp 中通过.Attr(scale).AttrType(OPTIONAL).Float(1.0f)与.Attr(softplus_threshold).AttrType(OPTIONAL).Float(20.0f)声明为可选默认值分别为 1.0 与 20.0tiling 阶段还会校验二者必须为有限值std::isfinite否则返回失败。约束说明使用 FusedGdnDecode 时必须满足以下约束64 K 2032Hv % H 0所有维度均为正数。stateRef 为 FLOAT32 时需满足ceil(K/16)*16-K 8保证 K 对齐到 16 时行填充字节数不超过 DataCopyPad 的 32 字节上限对应 tiling 中的MAX_DATA_COPY_PAD_BYTES检查。stateRef必须为连续 Tensor其他 Tensor 支持非连续输入。ssmStateIndices[i] 0表示无效槽位对应输出为 0 且stateRef不更新。ssmStateIndices[i] 0时用户需保证ssmStateIndices[i] BlockNum且同一批次内的正索引互不重复。aclnnFusedGdnDecode 默认确定性实现。关于 K 的上下界tiling 源码给出了更精确的推导MIN_SUPPORTED_K 64MAX_SUPPORTED_K由 uint8 索引上限与每 block 元素数推导而来见 op_host/fused_gdn_decode_tiling.cpp超出范围会直接报错返回。aclnn 接口调用说明FusedGdnDecode 的 aclnn 接口采用 CANN 标准的两段式调用先获取 workspace 与执行器再执行计算完整签名定义见 docs/aclnnFusedGdnDecode.md。函数原型第一段接口完成入参校验并返回 workspace 大小与执行器aclnnStatus aclnnFusedGdnDecodeGetWorkspaceSize( const aclTensor *mixedQkv, const aclTensor *a, const aclTensor *b, const aclTensor *aLog, const aclTensor *dtBias, aclTensor *stateRef, const aclTensor *ssmStateIndices, float scale, float softplusThreshold, aclTensor *out, uint64_t *workspaceSize, aclOpExecutor **executor);第二段接口在指定 Stream 上执行计算aclnnStatus aclnnFusedGdnDecode( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);参数使用说明第一段接口aclnnFusedGdnDecodeGetWorkspaceSize的核心入参语义如下参数名输入/输出描述使用说明mixedQkvaclTensor*输入按 Q、K、V 顺序拼接的输入不支持空 TensorBFLOAT16/FLOAT16ND 格式shape(B,2*H*KHv*V)支持非连续aaclTensor*输入门控输入 a不支持空 Tensor数据类型与 mixedQkv 一致shape(B,Hv)支持非连续baclTensor*输入门控输入 b不支持空 Tensor数据类型与 mixedQkv 一致shape(B,Hv)支持非连续aLogaclTensor*输入门控参数 A不支持空 TensorFLOAT32shape(Hv,)支持非连续dtBiasaclTensor*输入softplus 偏置不支持空 Tensor数据类型与 mixedQkv 一致shape(Hv,)支持非连续stateRefaclTensor*输入输出循环状态矩阵不支持空 Tensor必须为连续 TensorFLOAT32/BFLOAT16/FLOAT16shape(BlockNum,Hv,V,K)ssmStateIndicesaclTensor*输入batch 到 stateRef 槽位的映射不支持空 TensorINT32shape(B,)支持非连续scalefloat输入Q 归一化后的缩放系数必须为有限值softplusThresholdfloat输入softplus 阈值必须为有限值outaclTensor*输出GDN 输出不支持空 Tensor数据类型与 mixedQkv 一致shape(B,1,Hv,V)支持非连续workspaceSizeuint64_t*输出返回 Device 侧 workspace 大小-executoraclOpExecutor**输出返回 op 执行器-返回值aclnnStatus的具体状态码可参考 docs/zh/context/aclnn_return_code.md。第一段接口的入参校验失败场景包括返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001必选 Tensor、workspaceSize 或 executor 为空指针ACLNN_ERR_PARAM_INVALID161002输入数据类型、shape、连续性或属性不满足约束最小调用示例#include aclnnop/aclnn_fused_gdn_decode.h uint64_t workspaceSize 0; aclOpExecutor *executor nullptr; aclnnStatus ret aclnnFusedGdnDecodeGetWorkspaceSize( mixedQkv, a, b, aLog, dtBias, stateRef, ssmStateIndices, scale, softplusThreshold, out, workspaceSize, executor); if (ret ACLNN_SUCCESS) { ret aclnnFusedGdnDecode(workspace, workspaceSize, executor, stream); }完整可运行样例剖析仓库提供了完整可编译的示例 examples/test_aclnn_fused_gdn_decode.cpp其关键流程值得逐一对照ACL 环境初始化aclInit→aclrtSetDevice→aclrtCreateContext→aclrtSetCurrentContext→aclrtCreateStream参数构造示例使用 FLOAT16 输入 FLOAT32 状态batch 1qkHeads 1valueHeads 2keyDim 64valueDim 8stateSlots 2scale 0.125fsoftplusThreshold 20.0fmixedQkv shape 为{1, 2*1*64 2*8} {1, 144}stateRef shape 为{2, 2, 8, 64}out shape 为{1, 1, 2, 8}张量创建模板函数CreateTensor先aclrtMalloc分配 Device 内存再aclrtMemcpy将 Host 数据搬入最后用aclCreateTensor以ACL_FORMAT_ND创建 aclTensor两段式调用先调aclnnFusedGdnDecodeGetWorkspaceSize获取workspaceSize与executor按需aclrtMallocworkspace再调aclnnFusedGdnDecode执行结果回读aclrtSynchronizeStream同步后aclrtMemcpy将 out 从 Device 拷贝回 Host 并打印前 8 个元素资源释放依次释放 workspace、tensor、Device 内存、Stream、Context 并aclFinalize。该示例中stateIndicesData {1}即 batch 0 映射到 stateRef 的第 1 个槽位正索引执行后该槽位状态被原地更新。具体编译与运行方式可参考 docs/zh/context/compile_and_run_sample.md。torch 接口调用说明除 aclnn 接口外仓库还提供 PyTorch 扩展封装见 torch_ops_extension/README.md安装后可通过torch.ops.custom.npu_fused_gdn_decode调用torch.ops.custom.npu_fused_gdn_decode( mixed_qkv, a, b, a_log, dt_bias, state_ref, ssm_state_indices, scale, softplus_threshold20.0, )state_ref为原地更新 Tensor对应 aclnn 接口中的 stateRef输出 Tensor shape 为[B, 1, Hv, V]softplus_threshold为关键字参数默认 20.0。编译安装编译安装前需先安装 FusedGdnDecode 自定义算子包并配置自定义算子包环境变量然后执行cd experimental/attention/fused_gdn_decode/torch_ops_extension bash build_and_install.sh安装完成后导入custom_ops即可完成torch.ops.custom注册import torch import torch_npu import custom_opstorch.compile 成图支持torch 扩展还通过 custom_ops/converter/npu_fused_gdn_decode.py 为torch.ops.custom.npu_fused_gdn_decode.default注册了 fx2ge converterregister_fx_node_ge_converter使该自定义算子可被torchair在 torch.compile 场景下转换为 GE 图节点FusedGdnDecode图中输出为[out, state_out]其中state_out即原地更新的状态。这意味着该算子既支持 eager 模式直接调用也支持编译期成图后接入既有计算图。源码级实现原理算子注册与 shape 推导算子定义 op_host/fused_gdn_decode_def.cpp 声明了 7 个输入mixed_qkv、a、b、a_log、dt_bias、state、ssm_state_indices与 2 个输出out、state_out并声明DynamicShapeSupportFlag(true)等动态能力表明算子支持动态 shapeshape 推导 op_host/fused_gdn_decode_infershape.cpp 根据 mixedQkv 的 batch 维与 stateRef 的 Hv、V 维推导 out 的 shape(B, 1, Hv, V)state_out 与 state 同形数据类型分别跟随 mixedQkv 与 stateRef。tiling 策略tiling 实现 op_host/fused_gdn_decode_tiling.cpp 承担 shape 校验与切分策略决策核心要点包括输入校验校验各输入 rank 与维度为正、mixedQkv 与 H/K/Hv/V 的一致性、K 取值范围、dtype 组合等任一不满足即返回失败dtype 组合与 tilingKey根据 mixed 与 state 的 dtype 组合生成 4 种 tilingKey——BF16FP32、FP16FP32、BF16BF16、FP16FP16kernel 端据此实例化不同的模板参数任务切分以batch * hv为总任务数按 AIV 核数均分到各核并推导每个 block 所需的状态索引缓冲区大小BV 分块选择在 V 维上按候选集 {128, 64, 32, 16, 8} 从大到小尝试分块记为 bv综合估算 Q/K 缓冲、状态槽缓冲、标量 UB 与临时矩阵的占用选取能装入统一缓冲区UB的最大 bv保证 stateRef 为 FP32 时无需额外 FP32 计算副本而 BF16/FP16 状态则需要额外的 FP32 计算缓冲EstimateComputeBytes中computeMatrixBytes项的差异输出 tiling 数据将 batch、h、hv、k、v、bv、vTiles、各 stride、scale、softplusThreshold 等写入FusedGdnDecodeTilingData供 kernel 使用。kernel 执行流水kernel 入口 op_kernel/fused_gdn_decode.cpp 按 tilingKey 实例化KernelFusedGdnDecodeInType, StateType, STATE_FP32AIV 核执行核心实现位于 op_kernel/fused_gdn_decode.h主要流程为任务分配与状态索引加载每个 AIV 核按任务区间加载本核负责 batch 的ssmStateIndices到 UBQ/K 归一化Normalize通过平方求和、WholeReduceSum、加 ε、Rsqrt得到 RMS 归一化系数再批量乘回Q 额外乘以 scale门控计算PrepareGatingGroup在标量 UB 中完成 softplus基于softplusThreshold的数值稳定实现、-exp(A)、exp(g)与sigmoid(beta)的求值状态更新与输出对每个 V 头分块先按exp(g)缩放状态再经MatVecMul计算S_jk_h、ReduceRows得到delta beta*(v - exp(g)*S_jk_h)通过RankOneUpdate完成delta·k_h^T秩一更新写回状态最后MatVecMul(S_j, q_h)与ReduceRows得到输出状态更新结果通过stateOutQueue流水写回stateRef对应槽位无效槽位处理ssmStateIndices[i] 0时跳过计算直接Duplicate(0)写出零输出且不触碰 stateRef与约束说明一致。整条流水使用TQue/TBuf双缓冲BUFFER_NUM 1 的乒乓队列隐藏 GM 与 UB 间的搬移开销并以KERNEL_TYPE_AIV_ONLY指定纯向量核执行模式。总结与使用建议FusedGdnDecode 是 GDN 类线性注意力模型在 NPU 上做自回归解码的高效实现通过一次算子调用完成 QKV 拆分、归一化、门控、状态递推与输出投影避免了逐算子调度带来的多次 kernel 启动与中间结果搬移开销。集成时建议重点把握以下几点平台匹配仅 Atlas A2/A3 训练与推理系列产品支持使用前先确认目标机型shape 规划严格遵循64 K 2032、Hv % H 0与 stateRef 连续等约束FP32 状态还需满足 K 对齐的填充限制状态槽位管理利用ssmStateIndices在多个 batch/请求间复用 stateRef 槽位正索引不可重复、越界非正索引视为无效请求输出 0、不更新状态适合 KV 复用与批处理场景接口选择C 系列算子开发可走两段式 aclnn 接口PyTorch 训练/推理脚本可直接安装 torch_ops_extension 后调用torch.ops.custom.npu_fused_gdn_decode并可配合 torch.compile 成图深入定位问题遇到 shape 不合法或 UB 不足等问题时可结合 op_host/fused_gdn_decode_tiling.cpp 中对应的OP_LOGE错误日志定位根因。如需进一步理解两段式接口与编译运行细节可参阅 docs/zh/context/basic_concept.md 与 docs/zh/context/compile_and_run_sample.md算子的 shape 与 tiling 单测位于 tests/ut/op_host/可作为验证与二次开发的参考。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价