ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

CANN ops-transformer 中 ApplyRotaryPosEmbGrad 算子详解:双路旋转位置编码反向计算的融合实现与调用指南

CANN ops-transformer 中 ApplyRotaryPosEmbGrad 算子详解:双路旋转位置编码反向计算的融合实现与调用指南 CANN ops-transformer 中 ApplyRotaryPosEmbGrad 算子详解双路旋转位置编码反向计算的融合实现与调用指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerApplyRotaryPosEmbGrad 是 CANN ops-transformer 算子库中旋转位置编码RoPERotary Position Embedding系列的反向算子它将 query 与 key 两路的 RoPE 梯度计算融合进一次 kernel 调用同时可选输出 cos/sin 的梯度。本文以 README 为主体结合仓库内 aclnn 接口文档、PyTorch 封装、tiling 与 kernel 源码完整讲解该算子的数学原理、参数约束、两种调用方式以及底层多模板调度实现帮助你直接上手训练场景下的 RoPE 反向计算。一、算子功能与应用场景该算子是双路旋转位置编码算子 ApplyRotaryPosEmb 的反向算子核心功能如下执行双路反向计算同时计算query和key的 rope 反向梯度融合为一次 kernel 调用节省开销相比分别对 query、key 各执行一次反向 kernel融合实现节省了 cos/sin 的重复加载和 kernel launch 开销可选计算 cos/sin 梯度当正向输入query、key同时传入时额外计算grad_cos与grad_sin供需要更新 cos/sin 参与反传的场景使用。从产品支持情况看该算子目前仅面向Ascend 950PR / Ascend 950DT产品即源码配置目录 config/ascend950 对应的 SoCAtlas A2/A3、Atlas 200I/500 A2、Atlas 推理/训练系列等产品均不支持。表格如下产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品×Atlas A2 训练系列产品/Atlas A2 推理系列产品×Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×二、数学原理与计算公式旋转位置编码RoPE的核心思想是在 Head-Dim 维度上将每个头的向量按 D/2 拆成前后两半通过 cos/sin 旋转矩阵完成位置信息注入。反向计算即对该旋转过程求导。取正向计算中cos、sin发生 broadcast 的轴列表为dims即 cos/sin 中取值为 1、而grad_query_embed/grad_key_embed中对应维度大于 1 的轴包含 N 轴以及 BSND、SBND 布局下可选的 B 轴rotary_mode为half时的计算公式如下。1对输入梯度做末维对半切分$$ grad_q_1, grad_q_2 chunk(grad_query_embed, chunks2, dim-1) $$$$ grad_k_1, grad_k_2 chunk(grad_key_embed, chunks2, dim-1) $$$$ cos_1, cos_2 chunk(cos, chunks2, dim-1) $$$$ sin_1, sin_2 chunk(sin, chunks2, dim-1) $$2构造旋转后的正向向量用于 cos/sin 梯度$$ query_rotate cat((-query_2, query_1), dim-1) $$$$ key_rotate cat((-key_2, key_1), dim-1) $$3计算 query/key 的梯度$$ grad_query cat(cos_1 * grad_q_1 sin_2 * grad_q_2, cos_2 * grad_q_2 - sin_1 * grad_q_1, dim-1) $$$$ grad_key cat(cos_1 * grad_k_1 sin_2 * grad_k_2, cos_2 * grad_k_2 - sin_1 * grad_k_1, dim-1) $$4当同时传入 query 和 key 时沿广播轴 dims 归约得到 cos/sin 梯度$$ grad_cos sum(grad_query_embed * query grad_key_embed * key, dims) $$$$ grad_sin sum(grad_query_embed * query_rotate grad_key_embed * key_rotate, dims) $$仓库中的测试 golden 脚本 tests/assets/golden.py 以注释形式完整复现了上述公式并注明“所有路径统一升 FP32 计算结果转回输入 dtype”可供理解参考。三、参数说明各参数的完整说明如下表参数名输入/输出/属性描述数据类型数据格式grad_query_embed输入正向输出 query 的导数对应公式中 $grad_q_{embed}$。BFLOAT16、FLOAT16、FLOAT32NDgrad_key_embed输入正向输出 key 的导数对应公式中 $grad_k_{embed}$。BFLOAT16、FLOAT16、FLOAT32NDcos输入正向计算输入 cos需与 grad_query_embed 数据类型一致。BFLOAT16、FLOAT16、FLOAT32NDsin输入正向计算输入 sin需与 grad_query_embed 数据类型一致。BFLOAT16、FLOAT16、FLOAT32NDquery可选输入正向计算输入 query。如果为空指针则不计算 grad_cos 和 grad_sin必须与 key 同时传入或同时不传入。BFLOAT16、FLOAT16、FLOAT32NDkey可选输入正向计算输入 key。如果为空指针则不计算 grad_cos 和 grad_sin必须与 query 同时传入或同时不传入。BFLOAT16、FLOAT16、FLOAT32NDrotary_mode属性旋转模式仅支持 half。STRING-layout属性输入 Tensor 的布局格式。1BSND2SBND4TND。默认值为 1。INT64-grad_query输出正向计算输入 query 的导数shape 与 grad_query_embed 相同。BFLOAT16、FLOAT16、FLOAT32NDgrad_key输出正向计算输入 key 的导数shape 与 grad_key_embed 相同。BFLOAT16、FLOAT16、FLOAT32NDgrad_cos输出正向计算输入 cos 的导数仅当 query 和 key 非空时有效。BFLOAT16、FLOAT16、FLOAT32NDgrad_sin输出正向计算输入 sin 的导数仅当 query 和 key 非空时有效。BFLOAT16、FLOAT16、FLOAT32ND关于 layout 的补充说明BBatch批量大小SSeq-Length序列长度NHead-Num多头数DHead-Dim每个头的隐藏维度大小TB 和 S 的合轴layout4时输入为 3 维 Tensor其他 layout 下为 4 维。从 算子定义源码 可以看到host 侧注册的输入输出 dtype 均为DT_FLOAT16 / DT_FLOAT / DT_BF16格式为FORMAT_ND属性默认值分别为rotary_modehalf、layout1与上述参数表完全对应同时注册了DynamicCompileStaticFlag / DynamicRankSupportFlag / DynamicShapeSupportFlag表明算子支持动态 shape。四、约束说明输入输出 Tensor 只支持 3 维或 4 维layout 为 1 或 2 时为 4 维layout 为 4 时为 3 维。输入输出 Tensor 的 dtype 必须相同。输入输出 Tensor 不支持空 Tensor各维度必须大于 0。输入输出 Tensor 的 layout 必须相同。输入输出 Tensor 的 D 轴必须相同在 half 模式下必须 ≤ 1024 且能被 2 整除。grad_query_embed、grad_query的 shape 必须相同grad_key_embed、grad_key的 shape 必须相同。对于任意 layoutgrad_query_embed和grad_key_embed除 N 维度外其它维度必须相同。cos、sin的 N 维度必须等于 1layout 为 1BSND或 2SBND时cos、sin的 B 维度可以等于 1也可以和grad_query_embed的 B 维度一致layout 为 4TND时cos、sin的 T 维度必须和grad_query_embed的 T 维度一致除 N及 BSND、SBND 布局下可选广播的 B维度外其余维度需与grad_query_embed一致。cos、sin、grad_cos、grad_sin的 shape 必须相同。query维度需与grad_query_embed一致key维度需与grad_key_embed一致且query和key必须同时传入或同时不传入。rotary_mode仅支持 half。layout仅支持 {1, 2, 4}对应 {BSND, SBND, TND}。3BNSD 为预留暂不支持。这些约束在 tiling 源码中均有对应的显式校验。例如 apply_rotary_pos_emb_grad_tiling.cpp 中CheckRotaryModeShapeRelation()校验 D 轴≤ 1024D_LIMIT且% 2 0HALF_MODE_COEFValidateBroadcastByLayout()按 BSND/SBND/TND 分别校验 cos/sin 的 B/S/T 与 N 轴广播关系TND 下 cos 的 T 轴必须等于 grad_query_embed 的 TN 轴必须为 1BSND/SBND 下 cos 的 B 轴必须为 1 或等于输入 BS 轴必须一致CheckShape()校验grad_query_embed与grad_key_embed除 N 轴4D 布局下的 dim 2外各维度相同CheckOptionalInput()校验 query 与 grad_query_embed、key 与 grad_key_embed、cos 与 grad_cos、sin 与 grad_sin 的 shape 全等。五、aclnn 调用方式两段式接口aclnn 调用遵循 CANN 单算子调用的两段式接口规范必须先调用第一段aclnnApplyRotaryPosEmbGradGetWorkspaceSize完成入参校验并计算 workspace 大小再调用第二段aclnnApplyRotaryPosEmbGrad执行计算。函数原型如下aclnnStatus aclnnApplyRotaryPosEmbGradGetWorkspaceSize( const aclTensor *gradQueryEmbed, const aclTensor *gradKeyEmbed, const aclTensor *cos, const aclTensor *sin, const aclTensor *queryOptional, const aclTensor *keyOptional, char *rotaryModeOptional, int64_t layout, const aclTensor *gradQueryOut, const aclTensor *gradKeyOut, const aclTensor *gradCosOut, const aclTensor *gradSinOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnApplyRotaryPosEmbGrad( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)5.1 第一段接口参数与返回值第一段接口的参数细节完整见 aclnnApplyRotaryPosEmbGrad 接口文档参数名输入/输出描述使用说明数据类型维度(shape)非连续TensorgradQueryEmbed输入正向输出 query 的导数对应 grad_q_embed不支持空 TensorBFLOAT16/FLOAT16/FLOAT324(layout 1/2)或 3(layout 4)√gradKeyEmbed输入正向输出 key 的导数对应 grad_k_embed与 gradQueryEmbed 类型和维度一致不支持空 Tensor同上同上√cos输入正向计算输入 cos与 gradQueryEmbed 类型和维度一致N 维必须为 1同上同上√sin输入正向计算输入 sin与 gradQueryEmbed 类型和维度一致N 维必须为 1同上同上√queryOptional可选输入正向输入 query空指针时不计算 gradCos/gradSin与 keyOptional 必须同时传入或同时不传同上同上√keyOptional可选输入正向输入 key空指针时不计算 gradCos/gradSin与 queryOptional 必须同时传入或同时不传同上同上√rotaryModeOptional输入旋转模式仅支持 halfSTRING--layout输入输入 Tensor 布局1-BSND2-SBND4-TND3-BNSD(预留)INT64--gradQueryOut输出query 的导数与 gradQueryEmbed 类型和维度一致同上同上×gradKeyOut输出key 的导数与 gradQueryEmbed 类型和维度一致同上同上×gradCosOut输出cos 的导数query/key 非空时有效与 gradQueryEmbed 类型和维度一致同上同上×gradSinOut输出sin 的导数query/key 非空时有效与 gradQueryEmbed 类型和维度一致同上同上×workspaceSize输出Device 侧需申请的 workspace 大小----executor输出op 执行器包含算子计算流程----第一段接口完成入参校验出现以下场景时报错具体错误码定义见 aclnn 返回码返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001必选输入 gradQueryEmbed、gradKeyEmbed、cos、sin 和必选输出 gradQueryOut、gradKeyOut 是空指针ACLNN_ERR_PARAM_INVALID161002输入输出数据类型/格式不在支持范围、shape 不满足校验、维度不在支持范围、queryOptional 与 keyOptional 未成对传入、或 rotaryMode/layout 不符合支持值第二段接口接收workspace、workspaceSize、executor与stream四个参数workspace为 Device 侧申请的临时内存地址workspaceSize由第一段接口计算得出executor为第一段接口返回的 op 执行器stream指定任务执行的 Stream 流。注意第二段接口不可重复调用每次执行需重新走两段式流程。5.2 完整调用示例仓库提供了可参考的完整示例 examples/test_aclnn_apply_rotary_pos_emb_grad.cpp核心流程如下编译与运行方法参考编译与运行样例#include acl/acl.h #include aclnnop/aclnn_apply_rotary_pos_emb_grad.h #include iostream #include vector // 1. 资源初始化aclInit / aclrtSetDevice / aclrtCreateStream固定写法 // 2. 构造输入输出以 BSND layout、D128 为例 std::vectorint64_t gradQEmbedShape {1, 1, 1, 128}; std::vectorint64_t gradKEmbedShape {1, 1, 1, 128}; std::vectorint64_t cosShape {1, 1, 1, 128}; // N 维必须为 1 std::vectorint64_t sinShape {1, 1, 1, 128}; std::vectorint64_t queryShape {1, 1, 1, 128}; std::vectorint64_t keyShape {1, 1, 1, 128}; std::vectorint64_t gradQueryOutShape {1, 1, 1, 128}; std::vectorint64_t gradKeyOutShape {1, 1, 1, 128}; std::vectorint64_t gradCosOutShape {1, 1, 1, 128}; std::vectorint64_t gradSinOutShape {1, 1, 1, 128}; int64_t layout 1; // BSND const char *rotaryModeOptional half; // 通过 aclrtMalloc aclrtMemcpy 将 host 数据搬入 device // 再以 aclCreateTensor(..., ACL_FORMAT_ND, ...) 构造各 aclTensor // 3. 第一段接口计算 workspace 大小并创建 executor uint64_t workspaceSize 0; aclOpExecutor *executor; ret aclnnApplyRotaryPosEmbGradGetWorkspaceSize( gradQueryEmbed, gradKeyEmbed, cos, sin, query, key, const_castchar *(rotaryModeOptional), layout, gradQueryOut, gradKeyOut, gradCosOut, gradSinOut, workspaceSize, executor); // 4. 按 workspaceSize 申请 device 内存 void *workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 5. 第二段接口执行计算 ret aclnnApplyRotaryPosEmbGrad(workspaceAddr, workspaceSize, executor, stream); // 6. 同步并取回结果 ret aclrtSynchronizeStream(stream); ret aclrtMemcpy(resultData.data(), ..., gradQueryOutDeviceAddr, ..., ACL_MEMCPY_DEVICE_TO_HOST); // 7. 释放 aclTensor、device 内存与 stream示例中 shape 均取{1, 1, 1, 128}BSN1、D128满足 D ≤ 1024 且可被 2 整除query/key 同时传入以触发 grad_cos/grad_sin 计算。实际使用时可根据模型配置替换为 BSND、SBND 或 TND 布局并注意 cos/sin 的 N 维必须为 1。六、PyTorch API 调用方式PyTorch 侧封装位于 torch_extension/apply_rotary_pos_emb_grad.py通过torch.library注册自定义算子并以PrivateUse1后端分发到 NPU。函数原型如下完整文档见 torchapi_apply_rotary_pos_emb_gradcann_ops_transformer.apply_rotary_pos_emb_grad( grad_query_embed, grad_key_embed, cos, sin, *, queryNone, keyNone, rotary_modehalf, layout1, ) - (Tensor, Tensor, Optional[Tensor], Optional[Tensor])6.1 参数说明参数名参数类型可选/必选描述数据类型维度(shape)grad_query_embedTensor必选正向输出query_out的梯度bfloat16、float16、float32layout4 时为 (T, Nq, D)其他 layout 下为 4 维 Tensorgrad_key_embedTensor必选正向输出key_out的梯度除 N 维度外 shape 需与 grad_query_embed 一致同 grad_query_embedlayout4 时为 (T, Nk, D)其他 layout 下为 4 维 TensorcosTensor必选正向计算输入的余弦值张量N 维度必须等于 1同 grad_query_embed与输入布局对应的 3 维或 4 维 TensorsinTensor必选正向计算输入的正弦值张量shape 需与 cos 一致同 grad_query_embed同 cosqueryTensor可选正向计算输入 query传入时计算 grad_cos 和 grad_sin必须与 key 同时传入或同时不传入默认 None同 grad_query_embed与 grad_query_embed 一致keyTensor可选正向计算输入 key传入时计算 grad_cos 和 grad_sin必须与 query 同时传入或同时不传入默认 None同 grad_query_embed与 grad_key_embed 一致rotary_modestr可选旋转编码模式仅支持 half默认 half--layoutint可选1 表示 BSND2 表示 SBND4 表示 TND默认 1--返回值说明grad_queryTensor正向输入 query 的梯度shape 和数据类型与 grad_query_embed 一致grad_keyTensor正向输入 key 的梯度shape 和数据类型与 grad_key_embed 一致grad_cosOptional[Tensor]正向输入 cos 的梯度query 和 key 均传入时 shape 与 cos 一致否则为Nonegrad_sinOptional[Tensor]正向输入 sin 的梯度query 和 key 均传入时 shape 与 sin 一致否则为None。封装源码中的_check_inputs函数在 Python 侧提前完成了与上节约束一致的校验dtype 仅支持 float16/float32/bfloat16 且必须一致、维度必须为 3 或 4、TND(4) 布局要求 3 维输入、query/key 必须成对出现、query 与 grad_query_embed shape 相等、key 与 grad_key_embed shape 相等、rotary_mode 仅支持 half 等。Meta 实现apply_rotary_pos_emb_grad_meta则负责 shape/dtype 推导支撑 Autograd 与 FakeTensor 场景。6.2 单算子模式调用示例import torch import torch_npu from cann_ops_transformer.ops import apply_rotary_pos_emb_grad torch_npu.npu.set_device(0) B 1 S 64 N 8 D 128 grad_query_embed torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) grad_key_embed torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) cos torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) # N 维为 1可广播 sin torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) query torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) key torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) grad_query, grad_key, grad_cos, grad_sin apply_rotary_pos_emb_grad( grad_query_embed, grad_key_embed, cos, sin, queryquery, keykey, rotary_modehalf, layout1, # BSND ) print(fgrad_query shape: {grad_query.shape}) print(fgrad_key shape: {grad_key.shape}) print(fgrad_cos shape: {grad_cos.shape}) print(fgrad_sin shape: {grad_sin.shape})该示例展示了 BSND 布局下带 cos/sin 梯度的完整调用cos/sin取(B, S, 1, D)N 维为 1 从而沿 N 轴广播若不需要 grad_cos/grad_sin将query、key均置为None即可返回的 grad_cos/grad_sin 为None。6.3 与正向算子的配套使用该算子为 apply_rotary_pos_emb 的反向算子。正向接口使用rotary_modehalf时对 loss 执行.backward()会自动触发本算子仅在需要显式控制梯度时才需要手动调用本接口。该接口支持训练场景下单算子模式调用且默认支持确定性计算aclnn 与 torch API 文档均明确标注“默认确定性实现”。七、底层实现从算子定义到多模板调度7.1 算子定义与 shape 推导apply_rotary_pos_emb_grad_def.cpp注册 6 个输入grad_query_embed、grad_key_embed、cos、sin 为 REQUIREDquery、key 为 OPTIONAL、4 个输出grad_query、grad_key 为 REQUIREDgrad_cos、grad_sin 为 OPTIONALdtype 支持 FLOAT16/FLOAT/BF16格式统一 ND并声明DynamicCompileStaticFlag(true)、DynamicRankSupportFlag(true)、DynamicShapeSupportFlag(true)AICore 配置仅注册 ascend950。apply_rotary_pos_emb_grad_infershape.cppgrad_query 继承 grad_query_embed 的 shape、grad_key 继承 grad_key_embed 的 shape、grad_cos 继承 cos 的 shape、grad_sin 继承 sin 的 shapeSetGradOutputShapedtype 同理逐输出透传动态 shape 下以 -2unknown rank/-1unknown dim占位。7.2 tiling 阶段的广播判定与三模板调度tiling 核心逻辑位于 apply_rotary_pos_emb_grad_tiling.cpp 及三个模板实现文件_bab/_ab/_a。host 侧先执行一整套参数校验dtype、维度、D 轴上限与奇偶、cos/sin 广播关系、query/key 成对性等随后按 shape 关系判定内部ApplyRopeGradLayoutTND 布局3 维输入退化为 B1 的 BSND 处理若 N1 则走 NO_BROADCASTA 模板BSND 布局cos 的 B 轴等于 1 时进入 BSND 广播BAB 模板cos 的 B 轴与输入 B 一致时进入 SBNDAB 模板若 cos 与输入各维度完全一致含 gk 的 N 轴则判定为无广播A 模板SBND 布局shape 完全一致时回落 NO_BROADCASTA 模板否则按 AB 模板处理。最终通过三种 kernel tiling key 调度见 apply_rotary_pos_emb_grad_apt.cpp 中的枚举Tiling Key枚举值适用场景TILING_KEY_BAB203BSND 布局 cos B 轴1B/S/N 三层广播迭代TILING_KEY_AB204SBND 布局或 BSND 下 cos B 轴与输入一致TILING_KEY_A205无广播shape 完全一致最简路径7.3 kernel 执行流程与 workspacekernel 侧以__global__ __aicore__模板函数实现AIC 核直接返回、仅 AIV 核执行。计算分为两个 PhasePhase 1计算 grad_query / grad_key。BAB/AB 模板下若同时需要 grad_cos/grad_sinDcosFlag1会在 kernel 内同步累加 grad_cos/grad_sin 的部分积并写入 workspaceA 模板则分三个阶段先算 dx再预计算rotate(query)/rotate(key)写入 workspace最后通过ApplyDcosDsin高层 Mul 与 Q/K 累加直接写回 GM。Phase 2Reduce 归约。仅广播模板BAB/AB需要——将 Phase 1 产生的 dcos/dsin 部分积沿广播轴N 轴等跨核归约得到最终的 grad_cos/grad_sinA 模板因无广播无需 Reduce。workspace 大小由 tiling 计算并写入 tiling datausrWorkSpaceSize b * s * max(nQ, nK) * d * partialTypeSize * INPUT_OUTPUT_NUM两份部分积广播模板下 partialTypeSize 为 float 大小A 模板还有 16MB 的系统 workspace 预留这与第一段接口返回的workspaceSize直接对应。7.4 配置与测试验证算子二进制按 dtype 拆分为三个 binApplyRotaryPosEmbGrad_1/2/3对应 float16/bfloat16/float32见 apply_rotary_pos_emb_grad_binary.json单测覆盖 infershape 与 tiling 校验逻辑tests/ut/op_hostkernel 级 UT 直接包含_apt.cpp源码完成模板实例化覆盖 BAB/AB/A 三条路径tests/ut/op_kernel/arch35/test_apply_rotary_pos_emb_grad.cpp多路径 golden 脚本 tests/assets/golden.py 同时为 kernel spec、aclnn spec 与 E2E spec 提供参照实现q/k 两路分别归约以兼容 Nq ≠ Nk。八、总结ApplyRotaryPosEmbGrad 是 CANN ops-transformer 中面向 Ascend 950PR/950DT 的双路 RoPE 反向算子它一次 kernel 调用同时产出 grad_query、grad_key并在传入 query/key 时额外产出 grad_cos、grad_sin内部依据布局与广播关系在 BAB/AB/A 三种模板间自动选择配合 Reduce 归约与部分积 workspace 完成带广播的反向计算。实际使用时重点把握三类约束D 轴 ≤ 1024 且为偶数、cos/sin 的 N 维必须为 1、query 与 key 必须成对出现。相关接口文档与示例代码可直接参考 aclnnApplyRotaryPosEmbGrad 文档、torchapi 文档 与 调用示例。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表