CANN ops-transformer apply_rotary_pos_emb:NPU 上融合双路旋转位置编码的 PyTorch 算子全解析
CANN ops-transformer apply_rotary_pos_embNPU 上融合双路旋转位置编码的 PyTorch 算子全解析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本指南以 torchapi_apply_rotary_pos_emb.md 为骨架系统讲解 CANN ops-transformer 中apply_rotary_pos_emb算子旋转位置编码 RoPE的接口语义、三种旋转模式的计算公式、layout 布局约束、自动微分机制与 NPU 底层实现。读完本文你将能够在 Atlas 系列/Ascend 系列训练与推理场景下正确调用该算子完成 Q/K 双路位置编码的融合计算并理解其与反向算子apply_rotary_pos_emb_grad的联动关系以及 Tiling/Kernel 层的性能设计思路。一、算子概述一次 Kernel 调用完成 query 与 key 双路 RoPE在 Transformer 类大模型网络中旋转位置编码Rotary Position EmbeddingRoPE需要分别对注意力计算中的queryQ和keyK两路张量施加位置信息。传统实现通常将 Q、K 两路分别计算会带来额外的 kernel 启动开销与中间张量搬运成本。apply_rotary_pos_emb算子的核心设计目标即为提升网络性能将 query 和 key 两路旋转位置编码融合为一次 kernel 调用返回旋转位置编码后的 query 与 key 输出张量且输入张量不被修改输出为新分配张量非原地更新。该算子为 PyTorch 风格接口位于cann_ops_transformer.ops命名空间下底层封装了 ACLNN 层算子aclnnApplyRotaryPosEmb/aclnnApplyRotaryPosEmbV2同时支持单算子模式、图模式torch.compile torchair以及训练场景下的自动微分。二、产品支持情况产品是否支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品310P 等支持Atlas 训练系列产品910 等不支持注意从算子定义源码 apply_rotary_pos_emb_def.cpp 可以看到该算子在 AICore 侧实际注册了ascend910b、ascend910_93、ascend950、mc62、ascend310p、kirinx90、kirin9030等多个芯片配置覆盖 A2/A3/950 训练推理系列与 Kirin 处理器系列与文档中“Atlas 训练系列产品910不支持”的说明一致——910 原生配置未在 AICore 注册列表中。三、旋转位置编码原理与三种 rotary_mode 计算公式RoPE 的核心思想是在注意力计算前根据 token 位置对 Q、K 向量进行旋转使内积点积注意力自然携带相对位置信息。本算子通过rotary_mode属性控制旋转的配对方式支持half、quarter、interleave三种模式统一的计算形式为q_embed (query * cos) query_rotate * sin k_embed (key * cos) key_rotate * sin区别仅在于query_rotate/key_rotate的构造方式。以下公式在文档与 算子 README 中均有完整定义。3.1 half 模式默认最后一维二等分配对将query、key沿最后一维D 维即 Head-Dim二等分后执行旋转query_q1 query[..., : query.shape[-1] // 2] query_q2 query[..., query.shape[-1] // 2 :] query_rotate cat((-query_q2, query_q1), dim-1) q_embed (query * cos) query_rotate * sinkey的计算方式与query完全相同得到k_embed。该模式要求 D 维能被 2 整除。3.2 quarter 模式最后一维四等分两两配对将最后一维四等分前两段与后两段分别配对旋转query_q1 query[..., : D // 4] query_q2 query[..., D // 4 : D // 2] query_q3 query[..., D // 2 : D // 4 * 3] query_q4 query[..., D // 4 * 3 :] query_rotate cat((-query_q2, query_q1, -query_q4, query_q3), dim-1)该模式要求 D 维能被 4 整除。3.3 interleave 模式相邻两元素配对旋转取偶数下标元素与奇数下标元素分别作为一对query_q1 query[..., ::2].view(-1, 1) query_q2 query[..., 1::2].view(-1, 1) query_rotate cat((-query_q2, query_q1), dim-1).view(query.shape)该模式要求 D 维能被 2 整除。三种模式中half是默认值也是唯一支持自动微分的模式详见第五节。四、函数原型与参数说明4.1 函数原型cann_ops_transformer.apply_rotary_pos_emb(query, key, cos, sin, layoutBSND, rotary_modehalf) - (Tensor, Tensor)4.2 参数说明参数名参数类型可选/必选描述数据类型维度(shape)queryTensor必选待执行旋转位置编码的第一个张量。bfloat16、float16、float32layout为TND时3维其他layout下4维keyTensor必选待执行旋转位置编码的第二个张量。同query同querycosTensor必选旋转位置编码余弦值张量N维度必须等于1。同query同querysinTensor必选旋转位置编码正弦值张量shape需与cos一致。同query同coslayoutstr可选输入张量布局格式支持BSND、BSH、SBND、BNSD、TND。默认值为BSND。--rotary_modestr可选旋转编码模式支持half、quarter、interleave。默认值为half。--4.3 layout 语义说明layout为TND时输入为 3 维 Tensor其他layout下输入为 4 维 Tensor。其中 BBatch表示批量大小SSeq-Length表示序列长度NHead-Num表示多头数DHead-Dim表示每个头的隐藏维度大小T 表示 B 和 S 合轴常用于变长序列场景。BSH与BSND共用底层布局按BSND的维度语义处理。从 Python 封装 apply_rotary_pos_emb.py 可以看到底层 layout 的映射关系GRAD_LAYOUT_BY_LAYOUT { BSND: 1, BSH: 1, # BSH 映射为与 BSND 相同的整数 1 SBND: 2, BNSD: 3, TND: 4, }即BSH与BSND在底层共用 layout 数值 1按 BSND 的维度语义处理印证了文档中“共用底层布局”的说明。4.4 返回值说明query_outTensor旋转位置编码后的 query 输出张量为新分配张量shape 和数据类型与输入query一致输入query不被修改。key_outTensor旋转位置编码后的 key 输出张量为新分配张量shape 和数据类型与输入key一致输入key不被修改。这一“新分配输出、不修改输入”的语义在 C 封装 csrc/apply_rotary_pos_emb.cpp 中实现得非常直观先对输入执行query.clone()/key.clone()再以克隆结果为输出张量调用aclnnApplyRotaryPosEmbV2从而保证输入原张量不被改动。五、自动微分反向自动调用 apply_rotary_pos_emb_grad自动微分仅在Ascend 950PR / Ascend 950DT上支持。当query、key、cos、sin中任一输入requires_gradTrue且rotary_modehalf时该接口支持自动微分对 loss 执行.backward()时自动调用 apply_rotary_pos_emb_grad 计算query、key、cos、sin四路梯度。rotary_mode为quarter或interleave时不支持自动微分。从源码角度自动微分的实现位于 apply_rotary_pos_emb.py 的ApplyRotaryPosEmbFunction继承torch.autograd.Functionforward中调用底层 op 后通过ctx.save_for_backward(query, key, cos, sin)保存正向输入并记录layout与四个输入各自的needs_input_gradbackward中接收grad_query_embed、grad_key_embed随后调用torch.ops.cann_ops_transformer.apply_rotary_pos_emb_grad完成梯度计算rotary_mode被固定为half与文档“自动微分仅支持 half 模式”一致只有当cos或sin确实需要梯度时才会把query、key作为额外参数传入反向算子用于计算grad_cos、grad_sin否则传入None避免无谓计算顶层函数apply_rotary_pos_emb中当torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or cos.requires_grad or sin.requires_grad)时自动切换到ApplyRotaryPosEmbFunction.apply若此时rotary_mode ! half则直接抛出ValueError与文档约束完全对应。反向算子 apply_rotary_pos_emb_grad 同样将 query 与 key 两路梯度计算融合为一次 kernel 调用其layout参数为整数1 表示 BSND、2 表示 SBND、3 表示 BNSD、4 表示 TND且rotary_mode仅支持half。query与key必须同时传入或同时不传入两者均不传入时仅计算grad_query和grad_key返回的grad_cos和grad_sin为None。正常情况下开发者无需手动调用该反向接口只有需要显式控制梯度时才需直接使用。六、约束说明分产品6.1 通用约束该接口支持推理、训练场景下使用。该接口支持单算子模式和图模式调用。不支持空 Tensor。输入张量query、key、cos、sin的数据类型必须相同。cos、sin的 shape 必须相同且 N 维度必须等于 1。rotary_mode为 half 和 interleave 时输入 shape 最后一维D必须被 2 整除rotary_mode为 quarter 时输入 shape 最后一维D必须被 4 整除。6.2 Atlas 推理系列产品 / Atlas A2 / Atlas A3310P、910B、910_93layout仅支持 BSND 的 4 维 Tensor、TND 的 3 维 Tensor。layout为 BSND 时query、key、cos、sin输入 shape 的前 2 维B、S必须相等layout为 TND 时第 1 维T必须相等。query、key输入 shape 的最后一维D必须相等且等于 128 或 64cos、sin输入 shape 的最后一维D必须与之相等。这些检查在 Tiling 阶段被严格落地见 apply_rotary_pos_emb_tiling.cpp 中的CheckParamslayout属性仅允许 1BSND或 4TNDcos最后一维必须为 64 或 128query最后一维必须为 64 或 128 且不小于cos的最后一维cos的 N 轴必须为 1四个输入 dtype 必须一致。6.3 Ascend 950PR / Ascend 950DT训练场景自动微分仅支持rotary_modehalf。layout支持 BSND、SBND、BNSD 的 4 维 TensorTND 的 3 维 Tensor。对于任意layoutquery与key除 N 维度外其他维度必须相同。query、key输入 shape 的最后一维D必须相等且小于等于 1024cos、sin输入 shape 的最后一维D必须相等且小于等于query、key输入 shape 的最后一维D。layout为 BSND 时cos、sin的 B 维度可以等于 1也可以与query的 B 维度一致即支持沿 B 维广播。6.4 Atlas 推理系列产品310P不支持bfloat16。该限制同样适用于 Kirin X90 / Kirin 9030 处理器系列在 算子 README 中亦有说明算子定义源码中 310P/Kirin 系列的 AICore 配置仅注册了DT_FLOAT16与DT_FLOAT两种数据类型与文档一致。七、确定性计算默认支持确定性计算。八、底层实现原理源码级解读apply_rotary_pos_emb的完整调用链为PyTorch 前端 → torch 扩展 C 封装 → ACLNN 接口aclnnApplyRotaryPosEmb / aclnnApplyRotaryPosEmbV2→ 算子宿主侧InferShape / Tiling→ AICore Kernel。8.1 算子定义与 InferShape算子定义 apply_rotary_pos_emb_def.cpp声明 4 个必选输入query、key、cos、sin2 个输出同名query、key表示编码结果属性layoutint默认 1与rotary_modestring默认 half支持动态 Shape、动态 Rank 与动态编译。InferShape 实现 apply_rotary_pos_emb_infershape.cpp输出 shape 通过对query与cos/sin执行广播推导得到cos的 N 维为 1广播到query的 N 维并对动态 Shape未知维度-1与部分旋转位置编码cosD 维小于queryD 维做了专门处理。8.2 Tiling 设计按 BS 合轴切分 batch兼顾 UB 空间Tiling 逻辑见 apply_rotary_pos_emb_tiling.cpp其核心思路是将 BSND 的 B、S 两维或 TND 的 T 维合并为总 batch 数ab kDim0 * kDim1按可用 AIV 核数均分到各核尾核处理剩余 batchpreCoreBatch、useCoreNum、lastCoreBatch。依据数据类型与 UB 空间大小选择 TilingKey小 shape每核仅 1 个 batch 且 UB 放得下走TILINGKEY_SMALL否则按 A搬入/B计算乒乓流水拆分存在多次搬入一次搬出场景时走TILINGKEY_ABBF16 需要 cast 时走TILINGKEY_AB_CAST。一个 batch 的 UB 占用可近似为oneLoop qPart1Ub * 2 cosPart1Ub * 4 qPart1Ub q2q1Part1Ub * 2 isCast * (sin1UbSize * 2)其中qPart1Ub为搬运 (Q_n K_n) × D × castSize 的占用cosPart1Ub为搬运 1 × D × dtypeSize 的占用q2q1Part1Ub为旋转计算阶段的占用。Tiling 会计算 UB 一次可容纳的 batch 数据此推导外循环与内循环的次数分配preCBatchB、preCLTimes、lastCBatchL等并设置每核数据偏移qCoreOffset、kCoreOffset、cosCoreOffset。8.3 Kernel 实现AIV 核执行多 TilingKey 分派Kernel 入口见 op_kernel/apply_rotary_pos_emb.cpp声明为__global__ __aicore__的 AIV-only 任务根据 TilingKey 分派到不同的计算模板TILINGKEY_1SMALL→ARPESmallTILINGKEY_3AB→ARPEComputeABTILINGKEY_4AB_CAST→ARPEComputeABCast计算模板分别定义在 apply_rotary_pos_emb_small.h、apply_rotary_pos_emb_compute_ab.h 与 apply_rotary_pos_emb_compute_ab_cast.h 中。BF16 输入会在计算阶段 cast 为 float 以提升精度这正是 Tiling 中isCast分支与castDtypeSize的作用来源。8.4 图优化 Pass仓库还提供了图融合 pass apply_rotary_pos_emb_tensormove_pass.cpp用于在图模式torchair /torch.compile下对apply_rotary_pos_emb相邻的 tensor move 节点做优化进一步降低数据搬运开销。九、调用示例以下三个示例完整取自官方文档分别覆盖单算子模式、950 训练自动微分模式与图模式。9.1 单算子模式调用import torch import torch_npu from cann_ops_transformer.ops import apply_rotary_pos_emb torch_npu.npu.set_device(0) B 1 S 64 N 8 D 128 query torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) key torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) cos torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) sin torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) query_out, key_out apply_rotary_pos_emb( query, key, cos, sin, layoutBSND, rotary_modehalf, ) print(fOutput query shape: {query_out.shape}) print(fOutput key shape: {key_out.shape})注意cos、sin的 N 维取 1(B, S, 1, D)这是文档要求 N 维度必须等于 1 的直接体现同时 D 取 128满足 Atlas A2/A3/推理系列产品 D 必须为 64 或 128 的约束。9.2 Ascend 950 训练模式调用自动微分import torch import torch_npu from cann_ops_transformer.ops import apply_rotary_pos_emb torch_npu.npu.set_device(0) B, S, N, D 1, 64, 8, 128 query torch.randn(B, S, N, D, devicenpu, dtypetorch.float16, requires_gradTrue) key torch.randn(B, S, N, D, devicenpu, dtypetorch.float16, requires_gradTrue) cos torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16, requires_gradTrue) sin torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16, requires_gradTrue) # 正向返回q_embed、k_embed自动追踪计算图 query_out, key_out apply_rotary_pos_emb( query, key, cos, sin, layoutBSND, rotary_modehalf, # 自动微分仅支持half模式 ) loss query_out.sum() key_out.sum() loss.backward() # 自动调用apply_rotary_pos_emb_grad print(query.grad.shape) # query梯度 print(key.grad.shape) # key梯度 print(cos.grad.shape) # cos梯度 print(sin.grad.shape) # sin梯度四个输入全部置requires_gradTrue时backward()会通过ApplyRotaryPosEmbFunction.backward触发 apply_rotary_pos_emb_grad一次性产出四路梯度。9.3 图模式调用import torch import torch_npu import torchair from cann_ops_transformer.ops import apply_rotary_pos_emb torch_npu.npu.set_device(0) B, S, N, D 1, 64, 8, 128 class ApplyRotaryPosEmbModel(torch.nn.Module): def forward(self, query, key, cos, sin): return apply_rotary_pos_emb(query, key, cos, sin, layoutBSND, rotary_modehalf) model ApplyRotaryPosEmbModel().npu() npu_backend torchair.get_npu_backend() model torch.compile(model, backendnpu_backend, dynamicFalse) query torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) key torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) cos torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) sin torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) query_out, key_out model(query, key, cos, sin)图模式需要安装torchair组件并通过torchair.get_npu_backend()获取 NPU 后端接入torch.compile。十、配套接口与 ACLNN 调用方式除了 PyTorch 风格接口外该算子还提供两类 ACLNN 底层接口调用方式与对应示例见 算子 README调用方式调用样例说明aclnn 调用test_aclnn_apply_rotary_pos_emb.cpp通过aclnnApplyRotaryPosEmb接口方式调用 ApplyRotaryPosEmb 算子aclnn 调用test_aclnn_apply_rotary_pos_emb_v2.cpp通过aclnnApplyRotaryPosEmbV2接口方式调用PyTorch 封装底层实际使用的是 V2 接口ACLNN 接口的声明与实现分别位于 aclnn_apply_rotary_pos_emb_v2.h 与 aclnn_apply_rotary_pos_emb_v2.cpp采用标准的GetWorkspaceSizeaclnnXxx两段式调用模式。十一、相关资源索引正向接口文档torchapi_apply_rotary_pos_emb.md、aclnnApplyRotaryPosEmb.md、aclnnApplyRotaryPosEmbV2.md反向算子文档torchapi_apply_rotary_pos_emb_grad.mdPython 前端与自动微分封装torch_extension/apply_rotary_pos_emb.py、torch_extension/csrc/apply_rotary_pos_emb.cpp宿主侧实现apply_rotary_pos_emb_def.cpp、apply_rotary_pos_emb_infershape.cpp、apply_rotary_pos_emb_tiling.cppKernel 实现op_kernel/apply_rotary_pos_emb.cpp 及同目录下的apply_rotary_pos_emb_small.h、apply_rotary_pos_emb_compute_ab.h、apply_rotary_pos_emb_compute_ab_cast.h测试用例UT 侧 test_apply_rotary_pos_emb_infershape.cpp、test_apply_rotary_pos_emb_tiling.cpp、test_apply_rotary_pos_emb.cpp以及 ST 侧 executor_aclnnApplyRotaryPosEmb.py十二、总结apply_rotary_pos_emb是 CANN ops-transformer 中面向 Transformer 大模型推理/训练场景的高性能旋转位置编码算子其核心价值在于将 Q/K 两路 RoPE 计算融合为单次 kernel 调用并提供half/quarter/interleave三种旋转模式与BSND/BSH/SBND/BNSD/TND五种布局适配满足不同网络结构与推理引擎含变长序列 TND的需求。在 Ascend 950 上half模式还完整支持自动微分反向阶段自动衔接apply_rotary_pos_emb_grad使得该算子可以无缝嵌入端到端训练计算图。使用时请重点核对目标产品的 layout 支持范围、D 维取值64/128 或 ≤1024以及cos/sin的 N 维必须为 1 等约束。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考