ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

CANN ops-transformer 算子详解:aclnnRecurrentGatedDeltaRule 循环门控 Delta 规则算子

2026/9/19 11:32:44 拓冰建站 浏览量
CANN ops-transformer 算子详解:aclnnRecurrentGatedDeltaRule 循环门控 Delta 规则算子 CANN ops-transformer 算子详解aclnnRecurrentGatedDeltaRule 循环门控 Delta 规则算子【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer导读aclnnRecurrentGatedDeltaRule是 CANN ops-transformer 算子库中用于完成变步长variable-lengthRecurrent Gated Delta RuleRGDR计算的 AscendCL 单算子 API。RGDR 是一种将门控机制引入 Delta 规则Delta Rule的循环神经层更新方式在 Transformer 类大模型的线性注意力变体中承担隐藏状态SSM State的增量更新与注意力输出计算。阅读本文后你将掌握该算子的数学原理、产品支持范围、两段式接口用法、全部入参出参的数据类型与 shape 约束、调用示例的完整结构以及从源码层面理解其实现路径与测试验证方式。产品支持情况根据 aclnnRecurrentGatedDeltaRule.md 与 README.md该算子在以下产品上支持/不支持产品是否支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持该支持矩阵与算子定义中的 AICore 配置一一对应recurrent_gated_delta_rule_def.cpp 中为ascend910b对应 Atlas A2、ascend910_93对应 Atlas A3和ascend950分别注册了OpAICoreConfig未配置的平台则不支持。功能说明与数学原理算子功能接口功能完成变步长的 Recurrent Gated Delta Rule 计算。所谓变步长指 batch 内各序列的有效长度可以不同由actualSeqLengths描述算子按每个序列实际的 token 数量执行循环更新同时可通过ssmStateIndices将不同序列映射到独立的状态矩阵支持在预分配的 Page 化状态块中完成读写。计算公式Recurrent Gated Delta Rule循环门控 Delta 规则RGDR是一种应用于循环神经网络的算子也被应用于一种线性注意力机制中。在每个时间步 $t$网络根据当前的输入 $q_t$、$k_t$、$v_t$ 和上一个隐藏状态 $S_{t-1}$计算当前的注意力输出 $o_t$ 和新的隐藏状态 $S_t$。在这个过程中门控单元会决定有多少新信息存入隐藏状态以及有多少旧信息需要被遗忘。核心递推式与输出计算式如下$$ S_t : S_{t-1}\left(\alpha_t \mathrm{Diag}(\alpha_{kt})(I - \beta_t k_t k_t^T)\right) \beta_t v_t k_t^T \alpha_t \mathrm{Diag}(\alpha_{kt})S_{t-1} \beta_t (v_t - \alpha_t \mathrm{Diag}(\alpha_{kt})S_{t-1}k_t)k_t^T $$$$ o : \frac{S_t q_t}{\sqrt{d_k}} $$其中$S_{t-1}, S_t \in \mathbb{R}^{d_v \times d_k}$$q_t, k_t \in \mathbb{R}^{d_k}$$v_t \in \mathbb{R}^{d_v}$$\alpha_t \in \mathbb{R}$$\alpha_{kt} \in \mathbb{R}^{d_k}$$\beta_t \in \mathbb{R}$$o \in \mathbb{R}^{d_v}$。递推式第一项可理解为遗忘路径旧状态先被标量门控 $\alpha_t$ 与逐维门控 $\alpha_{kt}$ 衰减并减去沿 $k_t$ 方向被擦除的部分$\beta_t k_t k_t^T$ 项第二项为写入路径以 $\beta_t$ 为门控把 $v_t$ 与 $k_t$ 的外积增量写入状态。输出 $o$ 则是新状态与当前 query 的乘积并按 $1/\sqrt{d_k}$ 缩放。公式中的门控在 API 中并不是直接以 $\alpha$、$\beta$ 形式传入的而是$\alpha_t e^{g_t}$即g是衰减系数对数域$\alpha_{kt} e^{gk_t}$即gk是逐维衰减系数对数域$\beta_t$ 直接由beta张量传入。从源码看这一语义也被 kernel 层明确承接recurrent_gated_delta_rule.cpp 中的 kernel 入口recurrent_gated_delta_rule接收query, key, value, beta, state, cuSeqlens, ssmStateIndices, g, gk, numAcceptedTokens, out, stateOut全部入参并根据 tiling key 选择以 FP32 或 BF16 状态精度实例化的RGDR模板类执行Init Process。两段式接口与函数原型该算子遵循 CANN 单算子 API 的两段式接口规范详见 two_phase_api.md必须先调用第一段aclnnRecurrentGatedDeltaRuleGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器再申请 device 侧 workspace 内存然后调用第二段aclnnRecurrentGatedDeltaRule执行计算。第一段接口原型aclnnStatus aclnnRecurrentGatedDeltaRuleGetWorkspaceSize( const aclTensor *query, const aclTensor *key, const aclTensor *value, const aclTensor *beta, aclTensor *stateRef, const aclTensor *actualSeqLengths, const aclTensor *ssmStateIndices, const aclTensor *g, const aclTensor *gk, const aclTensor *numAcceptedTokens, float scaleValue, aclTensor *out, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口原型aclnnStatus aclnnRecurrentGatedDeltaRule( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)对应的头文件声明位于 aclnn_recurrent_gated_delta_rule.h其中以 Doxygen 注释标注了每个参数的 dtype 支持范围而 aclnn 层的实际编排实现位于 aclnn_recurrent_gated_delta_rule.cpp。注意第二段接口不能重复调用同一 executor 只能执行一次详见 two_phase_api.md。aclnnRecurrentGatedDeltaRuleGetWorkspaceSize 参数说明参数总表参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensorquery输入公式中的 q不支持空 TensorBFLOAT16ND(T, Nk, Dk)√key输入公式中的 k不支持空 TensorBFLOAT16ND(T, Nk, Dk)√value输入公式中的 v不支持空 TensorBFLOAT16ND(T, Nv, Dv)√beta输入公式中的 β不支持空 TensorBFLOAT16ND(T, Nv)√stateRef输入输出状态矩阵公式中的 S不支持空 TensorBFLOAT16、FLOAT32ND(BlockNum, Nv, Dv, Dk)√actualSeqLengths输入不同 batch 的有效序列长度不支持空 TensorINT32ND(B,)√ssmStateIndices输入输入序列到状态矩阵的映射索引不支持空 Tensorstate[ssmStateIndices[i]] 表示第 i 个 token 的状态矩阵INT32ND(T,)√g输入衰减系数公式中的 αe^g不支持空 Tensor若传入 nullptr 则表示全 0 的 tensorFLOAT32ND(T, Nv)√gk输入衰减系数公式中的 αke^gk不支持空 Tensor若传入 nullptr 则表示全 0 的 tensorFLOAT32ND(T, Nv, Dk)√numAcceptedTokens输入每个序列接受的 token 数量不支持空 TensorINT32ND(B,)√scaleValue输入query 的缩放因子对应公式中的 1/sqrt(d_k)-----out输出公式中的 o-BFLOAT16ND(T, Nv, Dv)√workspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----其中 $B$ 表示 batch size令 $L_i$ 表示第 i 个序列的长度则 $T\sum_{i}^{B} L_i$ 表示累积序列长度$N_k$ 表示 key 的头数$N_v$ 表示 value 的头数$D_k$ 表示 key 向量的维度$D_v$ 表示 value 向量的维度。参数设计要点解读结合算子定义源码 recurrent_gated_delta_rule_def.cpp可以进一步理解几个参数的底层设计可选输入OPTIONALg、gk、numAcceptedTokens在算子定义中均为OPTIONAL类型。当g/gk传 nullptr 时等价于全 0 tensor即 $\alpha\alpha_k1$无衰减tiling 阶段通过hasGama、hasGamaK、hasAcceptedTokens标志见 recurrent_gated_delta_rule_tiling_data.h区分可选输入是否真实存在从而避免多余访存。stateRef 双向语义stateRef既是输入又是输出对应算子定义中的stateInput 与stateOutputkernel 在读入旧状态后原地更新并写回这正是文档约束中仅支持 0 轴、1 轴非连续 Tensor的原因——非连续状态块如 Page 化 KV Cache 中的分块存储由 stride 描述见 TilingData 中的stateStride0、stateStride1字段。workspaceSize由第一段接口经 infershape 与 tiling 计算得到。tiling 类 recurrent_gated_delta_rule_tiling.h 定义了RecurrentGatedDeltaRuleTiling其职责链为获取平台信息核数、UB/L1/L0C 资源、解析 shape/attr/format、计算数据切分 TilingData、计算 TilingKey 并最终计算 workspace 大小。返回值与第一段接口校验aclnnStatus返回状态码具体参见 aclnn_return_code.md。第一段接口完成入参校验出现以下场景时报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, numAcceptedTokens, out 存在空指针。ACLNN_ERR_PARAM_INVALID161002输入 Tensor 的数据类型不在支持的范围内。ACLNN_ERR_PARAM_INVALID161002输入 Tensor 的数据格式不在支持范围内。ACLNN_ERR_PARAM_INVALID161002输入 Tensor 的 shape 不在支持范围内。除上述显式校验外第一段接口还会在内部完成输出 shape/数据类型推导InferShape。从 recurrent_gated_delta_rule_infershape.cpp 的实现可以看出out的 shape 完全继承自value的 (T, Nv, Dv)stateOut的 shape 完全继承自输入state的 (BlockNum, Nv, Dv, Dk)且out的 dtype 固定为 BF16stateOut的 dtype 与输入state保持一致BF16 或 FP32。这与 recurrent_gated_delta_rule.cpp 中AllocTensor以DT_BF16分配输出、并声明OP_OUTPUT(out, stateRef)双输出的 aclnn 编排逻辑一致。aclnnRecurrentGatedDeltaRule 参数说明第二段接口的参数较少全部为输入参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址。workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口 aclnnRecurrentGatedDeltaRuleGetWorkspaceSize 获取。executor输入op 执行器包含算子计算流程。stream输入指定执行任务的 Stream。返回值同样为aclnnStatus具体参见 aclnn_return_code.md。除上述通用返回码外执行阶段还可能返回 561xxx 系列内部异常如 tiling 异常、kernel 查找失败、OPP 路径未配置等排查时可使用aclGetRecentErrMsg获取详细错误信息。约束说明确定性计算aclnnRecurrentGatedDeltaRule 默认为确定性实现。这也是算子定义中为各平台配置ExtendCfgInfo(softsync.flag, true)的体现保证多次运行结果一致便于调试与精度复现。Shape 约束算子校验输入 shape 大小需满足$$0 L_i \le 8,\quad 0 N_k \le 256,\quad N_k \le N_v \le 256,\quad N_v\ %\ N_k 0$$$$0 D_k \le 512,\quad 0 D_v \le 512,\quad 0 T,\quad 0 B,\quad T \le BlockNum$$stateRef 仅支持 0 轴、1 轴非连续 Tensor。数值与语义约束用户需保证算子不校验以下约束由于算子无法获取 tensor 中的具体数值故需用户保证算子不校验$ssmStateIndices[i] BlockNum$$0 actualSeqLengths[i] \le 8$且 $actualSeqLengths[i]$ 累加和等于 $T$$1 \le numAcceptedTokens[i] \le actualSeqLengths[i]$$-1 \le query[i][j][k] \le 1$$-1 \le key[i][j][k] \le 1$$g[i][j] 0$保证 $\alpha e^g 1$即衰减$gk[i][j][k] 0$$0 beta[i][j] 1$从测试资产 assets/README.md 的描述可以印证actual_seq_lengths将 T 均分到 B 个 batch、ssm_state_indicesarange(T)、num_accepted_tokens每 batch 取[1, seq_len]内的值这些 int32 张量无法随机生成均由测试框架的inputs.py适配器按上述语义填充。调用示例以下示例代码来自文档仅供参考具体编译和执行过程请参考 compile_and_run_sample.md。仓库中另有可直接参考的完整可编译样例 test_aclnn_recurrent_gated_delta_rule.cpp 及 UT 版本 tests/ut/op_api/test_aclnn_recurrent_gated_delta_rule.cpp。#include iostream #include vector #include cstring #include acl/acl.h #include aclnnop/aclnn_recurrent_gated_delta_rule.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 shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } void PrintOutResult(std::vectorint64_t shape, void **deviceAddr, const char *name) { auto size GetShapeSize(shape); std::vectoruint16_t resultData(size, 0); auto ret aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(copy result from device to host failed. ERROR: %d\n, ret); return); for (int64_t i 0; i size; i) { if (i 5) { // print the first five data break; } float result 0.0f; uint32_t val static_castuint32_t(resultData[i]) 16U; std::memcpy(result, val, sizeof(result)); LOG_PRINT(%s result[%ld] is: %f\n, name, i, result); } } int Init(int32_t deviceId, aclrtContext *context, aclrtStream *stream) { // AscendCL初始化 auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed. ERROR: %d\n, ret); return ret); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); ret aclrtCreateContext(context, deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateContext failed. ERROR: %d\n, ret); return ret); ret aclrtSetCurrentContext(*context); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetCurrentContext failed. ERROR: %d\n, ret); return ret); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed. ERROR: %d\n, ret); return ret); return 0; } template typename T int CreateAclTensor(const std::vectorT hostData, const std::vectorint64_t shape, void **deviceAddr, aclDataType dataType, aclTensor **tensor) { auto size GetShapeSize(shape) * sizeof(T); // 调用aclrtMalloc申请device侧内存 auto ret aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed. ERROR: %d\n, ret); return ret); // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 ret aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMemcpy failed. ERROR: %d\n, ret); return ret); // 调用aclCreateTensor接口创建aclTensor *tensor aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return ACL_SUCCESS; } int main() { // 1.device/context/stream初始化参考AscendCL对外接口列表 // 根据自己的实际device填写deviceId int32_t deviceId 0; aclrtContext context; aclrtStream stream; auto ret Init(deviceId, context, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 构造输入与输出需要根据API的接口自定义构造 void *queryDeviceAddr nullptr; void *keyDeviceAddr nullptr; void *valueDeviceAddr nullptr; void *gDeviceAddr nullptr; void *betaDeviceAddr nullptr; void *stateRefDeviceAddr nullptr; void *actSeqLenDeviceAddr nullptr; void *ssmStaIdDeviceAddr nullptr; void *numAccTokDeviceAddr nullptr; void *attnOutDeviceAddr nullptr; aclTensor *query nullptr; aclTensor *key nullptr; aclTensor *value nullptr; aclTensor *g nullptr; aclTensor *gk nullptr; aclTensor *beta nullptr; aclTensor *stateRef nullptr; aclTensor *actSeqLen nullptr; aclTensor *ssmStaId nullptr; aclTensor *numAccTok nullptr; aclTensor *attnOut nullptr; // 自定义输入与属性 int32_t batchSize 2; int32_t mtp 2; int32_t headKNum 4; int32_t headVNum 8; int32_t dimV 32; int32_t dimK 32; std::vectorint64_t stateShape {batchSize * mtp, headVNum, dimV, dimK}; std::vectorint64_t qkShape {batchSize * mtp, headKNum, dimK}; std::vectorint64_t vShape {batchSize * mtp, headVNum, dimV}; std::vectorint64_t gShape {batchSize * mtp, headVNum}; std::vectorint64_t gkShape {batchSize * mtp, headVNum, dimK}; std::vectorint64_t actSeqLenShape {batchSize}; std::vectorint64_t ssmStaIdShape {batchSize * mtp}; std::vectorfloat stateRefHostData(GetShapeSize(stateShape)); std::vectorfloat queryHostData(GetShapeSize(qkShape)); std::vectorfloat keyHostData(GetShapeSize(qkShape)); std::vectorfloat valueHostData(GetShapeSize(vShape)); std::vectorfloat gHostData(GetShapeSize(gShape)); std::vectorfloat gkHostData(GetShapeSize(gkShape)); std::vectorfloat betaHostData(GetShapeSize(gShape)); std::vectorint32_t actSeqLenHostData(batchSize, mtp); std::vectorint32_t ssmStaIdHostData(batchSize * mtp); std::vectorint32_t numAccTokHostData(batchSize, 1); for (int i 0; i stateRefHostData.size(); i) { stateRefHostData[i] 0.5; } for (int i 0; i queryHostData.size(); i) { queryHostData[i] 0.5; } for (int i 0; i keyHostData.size(); i) { keyHostData[i] 0.5; } for (int i 0; i valueHostData.size(); i) { valueHostData[i] 0.5; } for (int i 0; i gHostData.size(); i) { gHostData[i] -0.5; } for (int i 0; i gkHostData.size(); i) { gkHostData[i] -0.5; } for (int i 0; i betaHostData.size(); i) { betaHostData[i] 0.5; } for (int i 0; i ssmStaIdHostData.size(); i) { ssmStaIdHostData[i] i; } std::vectorfloat attnOutHostData(valueHostData); ret CreateAclTensor(stateRefHostData, stateShape, stateRefDeviceAddr, aclDataType::ACL_BF16, stateRef); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(queryHostData, qkShape, queryDeviceAddr, aclDataType::ACL_BF16, query); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(keyHostData, qkShape, keyDeviceAddr, aclDataType::ACL_BF16, key); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(valueHostData, vShape, valueDeviceAddr, aclDataType::ACL_BF16, value); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gHostData, gShape, gDeviceAddr, aclDataType::ACL_FLOAT, g); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(gkHostData, gkShape, gkDeviceAddr, aclDataType::ACL_FLOAT, gk); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(betaHostData, gShape, betaDeviceAddr, aclDataType::ACL_BF16, beta); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(actSeqLenHostData, actSeqLenShape, actSeqLenDeviceAddr, aclDataType::ACL_INT32, actSeqLen); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(ssmStaIdHostData, ssmStaIdShape, ssmStaIdDeviceAddr, aclDataType::ACL_INT32, ssmStaId); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(numAccTokHostData, actSeqLenShape, numAccTokDeviceAddr, aclDataType::ACL_INT32, numAccTok); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(attnOutHostData, vShape, attnOutDeviceAddr, aclDataType::ACL_BF16, attnOut); CHECK_RET(ret ACL_SUCCESS, return ret); // 3. 调用CANN算子库API需要修改为具体的Api名称 uint64_t workspaceSize 0; float scale 1.0; aclOpExecutor *executor; // 调用aclnnRecurrentGatedDeltaRuleGetWorkspaceSize第一段接口 ret aclnnRecurrentGatedDeltaRuleGetWorkspaceSize(query, key, value, beta, stateRef, actSeqLen, ssmStaId, g, gk, numAccTok, scale, attnOut, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnRecurrentGatedDeltaRuleGetWorkspaceSize failed. ERROR: %d\n, ret); return ret); // 根据第一段接口计算出的workspaceSize申请device内存 void *workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(allocate workspace failed. ERROR: %d\n, ret); return ret); } // 调用aclnnRecurrentGatedDeltaRule第二段接口 ret aclnnRecurrentGatedDeltaRule(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnRecurrentGatedDeltaRule failed. ERROR: %d\n, ret); return ret); // 4. 同步等待任务执行结束 ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed. ERROR: %d\n, ret); return ret); // 5. 获取输出的值将device侧内存上的结果拷贝至host侧需要根据具体API的接口定义修改 PrintOutResult(stateShape, stateRefDeviceAddr, finalState); PrintOutResult(vShape, attnOutDeviceAddr, out); // 6. 释放aclTensor和aclScalar需要根据具体API的接口定义修改 aclDestroyTensor(query); aclDestroyTensor(key); aclDestroyTensor(value); aclDestroyTensor(g); aclDestroyTensor(gk); aclDestroyTensor(beta); aclDestroyTensor(stateRef); aclDestroyTensor(actSeqLen); aclDestroyTensor(ssmStaId); aclDestroyTensor(numAccTok); aclDestroyTensor(attnOut); // 7. 释放device资源 aclrtFree(queryDeviceAddr); aclrtFree(keyDeviceAddr); aclrtFree(valueDeviceAddr); aclrtFree(gDeviceAddr); aclrtFree(gkDeviceAddr); aclrtFree(betaDeviceAddr); aclrtFree(stateRefDeviceAddr); aclrtFree(actSeqLenDeviceAddr); aclrtFree(ssmStaIdDeviceAddr); aclrtFree(numAccTokDeviceAddr); aclrtFree(attnOutDeviceAddr); if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtDestroyContext(context); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例关键点解读shape 设计示例采用batchSize2, mtp2因此 $T batchSize \times mtp 4$且actSeqLenHostData每个 batch 均填mtp2满足有效序列长度累加和等于 T的约束ssmStaIdHostData[i] i使第 i 个 token 映射到第 i 个状态块配合stateShape (4, Nv, Dv, Dk)满足 $T \le BlockNum$。数值范围query/key/value 取 0.5在 $[-1, 1]$ 内g/gk 取 -0.5满足 $ 0$beta 取 0.5满足 $(0, 1)$与用户需保证的数值约束一一对应。BF16 解码PrintOutResult中uint16_t数据通过左移 16 位再 memcpy 为 float 的方式完成 BF16→FP32 解码打印。内存生命周期按申请→拷贝→创建 aclTensor→调用两段接口→同步→回拷结果→销毁 tensor→释放 device 内存的标准流程组织其中 workspace 仅在workspaceSize 0时才申请与释放。源码级实现路径算子注册与平台配置recurrent_gated_delta_rule_def.cpp 通过OpDef注册算子10 个输入g/gk/num_accepted_tokens为 OPTIONAL、2 个输出out与state、1 个属性scale_valueOPTIONAL默认 1.0。平台相关配置中ascend910b与ascend910_93使用默认 kernel 文件ascend950额外指定opFile.value recurrent_gated_delta_rule_apt对应 recurrent_gated_delta_rule_apt.cpp 的高阶 API 实现。InferShape 与输出推导recurrent_gated_delta_rule_infershape.cpp 定义了 shape 与 dtype 推导规则out直接继承value的三维 shapestateOut继承输入state的四维 shapedtype 方面out恒为DT_BF16stateOut与输入state同 dtype。Tiling 与状态精度分支tiling 类 recurrent_gated_delta_rule_tiling.h 完成平台资源获取、shape/attr 解析、可选输入探测AnalyzeOptionalShapes/GetOptionalInput、UB 容量计算CalUbSize、scale 获取GetScale以及 workspace 计算TilingData 结构见 recurrent_gated_delta_rule_tiling_data.h记录了核数、UB 容量、T/Nk/Dk/Nv/Dv 等核心维度、scale、可选输入标志以及非连续 state 的 stride。kernel 入口 recurrent_gated_delta_rule.cpp 依据 recurrent_gated_delta_rule_tiling_key.h 中定义的两个 TilingKey 分支TILING_KEY_RGDR_FP32_STATEstate 以 FP32 累积kernel 模板参数为RGDRbfloat16_t, bfloat16_t, float与TILING_KEY_RGDR_BF16_STATEstate 以 BF16 存储RGDRbfloat16_t, bfloat16_t, bfloat16_t。具体循环逻辑含 AIV 矢量核上的矩阵-向量运算在 arch22/recurrent_gated_delta_rule_arch22.h 与 arch35/recurrent_gated_delta_rule_arch35.h 中按架构实现。测试与精度验证仓库为算子提供了多层测试UT 单测op_api、op_hostinfershape、tiling、op_kernel三层 UT见 tests/ut 目录。pytest 精度用例tests/pytest 提供参数化用例与 CPU golden 参考实现 recurrent_gated_delta_rule_golden.py。TTK 测试资产tests/assets/README.md 说明基于 ops-test-kitTTK框架支持三种验证方式E2E eagertorch_npu.npu_recurrent_gated_delta_rule直调、E2E 静态图torch.compile、ACLNN 模式本接口。精度判定采用np.isclose(rtol0.0078125, atol0.0001)bfloat16 社区标准相对误差通过率 ≥ 99.5% 且最大相对误差 10.0 方为 Pass且对output state双输出同时比较。其中 ACLNN 模式下 inputs/golden 适配器按aclnnRecurrentGatedDeltaRuleGetWorkspaceSize的实参顺序对齐convert_rdv_to_csv.py可将 148 条 RDV 全量用例转换为 TTK CSV覆盖连续/非连续 state、bf16/fp32 state dtype 组合。常见问题与注意事项两段式接口缺一不可直接调用第二段接口或重复调用第二段接口均不符合规范必须先由第一段接口取得workspaceSize与executor。workspaceSize 为 0 的情况示例中workspaceSize 0时才执行aclrtMalloc申请与释放需与第一段接口返回值保持一致。用户自保证的数值语义ssmStateIndices越界、actualSeqLengths累加和与 T 不一致、g/gk 非负、beta 越界等场景算子不校验可能导致错误结果或设备异常务必在业务侧先行校验。stateRef 非连续限制仅支持 0 轴、1 轴非连续这是 Page 化状态存储的核心能力其他轴非连续场景未在支持范围内。确定性算子为确定性实现结果可复现若在测试中发现精度抖动应优先排查输入数值是否落在上述用户自保证的约束区间内。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考