CANN ops-transformer 算子详解:inplace_partial_rotary_mul 原地部分旋转位置编码
2026/9/20 17:04:24 网站建设 项目流程
  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

本篇文章围绕 CANN ops-transformer 开源仓库中的inplace_partial_rotary_mul算子展开,系统讲解其在 NPU 上以 Inplace 方式执行单路旋转位置编码(RoPE)的接口设计、interleave 旋转计算公式、partial_slice局部切片语义、自动微分链路,以及单算子模式、训练模式与图模式三种调用方式。读完本文,你将能够正确配置该算子的输入与约束,理解其与反向算子inplace_partial_rotary_mul_backward的联动机制,并直接复用文中完整示例完成 NPU 上的 RoPE 编码。

算子概览与产品支持情况

inplace_partial_rotary_mul是 CANN transformer 类大模型算子库(项目主页)中 posembedding 目录下的核心位置编码算子,源码位于 posembedding/inplace_partial_rotary_mul。其最显著的特点是原地(Inplace)计算:执行单路旋转位置编码时直接修改输入张量x,不产生新的输出张量;同时支持通过partial_slice参数指定输入张量最后一维上的局部范围,仅对该范围内的数据执行旋转位置编码,其余位置保持原值。这一设计特别适合在多头注意力模型中对隐藏维度做"部分旋转"(如仅旋转前 D/2 维)的常见 RoPE 用法,可显著节省显存与访存开销。

该算子在不同 NPU 产品上的支持情况如下(依据 torchapi_inplace_partial_rotary_mul.md):

产品支持情况
Ascend 950PR / Ascend 950DT支持
Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持
Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持
Atlas 200I/500 A2 推理产品不支持
Atlas 推理系列产品不支持
Atlas 训练系列产品不支持

从算子注册代码 inplace_partial_rotary_mul_def.cpp 可以看到,底层 AICore 侧分别针对ascend910bascend910_93(即 Atlas A2/A3 系列)与ascend950配置了不同的算子实现文件,其中 ascend950 走inplace_partial_rotary_mul_apt实现路径,与上文产品支持矩阵一一对应。

功能与计算公式

接口功能

执行单路旋转位置编码的 Inplace 计算,直接修改输入张量x,不产生新的输出张量。该接口支持通过partial_slice参数指定输入张量最后一维上的局部范围,仅对该范围内的数据执行旋转位置编码,其余位置保持原值。

输入x采用BSND维度格式:

  • B(Batch):批量大小;
  • S(Seq-Length):序列长度;
  • N(Head-Num):多头数;
  • D(Head-Dim):每个头的隐藏维度大小。

partial_slice作用于输入x的最后一维 D 维,取值范围为左闭右开区间[start, end)。Python 接口中partial_slice默认值为None,内部按[0, 0]处理(对应"空切片、不执行旋转"的语义,见 inplace_partial_rotary_mul.py)。

计算公式

interleave 模式(rotary_mode"interleave")下,设partial_slice=[start, end],被旋转的局部张量为:

$$x_{slice} = x[..., start:end]$$

计算过程如下:

$$x_1 = x_{slice}[..., ::2]$$

$$x_2 = x_{slice}[..., 1::2]$$

$$x_{rotate} = \text{cat}(-x_2, x_1)$$

$$x_{out} = x_{slice} \cdot \cos + x_{rotate} \cdot \sin$$

最终将 $x_{out}$ 原地写回x[..., start:end]。其中,$x$ 表示参数x,$\cos$ 表示参数r1,$\sin$ 表示参数r2。当startend相等时,不执行旋转位置编码,x保持不变。反向算子inplace_partial_rotary_mul_backward同样支持空 Tensor 和切片长度为零的场景(执行 no-op)。

从底层 Kernel 实现可以印证上述公式:在 inplace_partial_rotary_mul.h 的Process中,先通过SetGatherSrcOffset构造奇偶索引交换的 Gather 表(idsUb.SetValue(i, i ^ 1),即以 XOR 1 交换偶数位与奇数位),再依次完成x * cosComputeMul)、奇偶交换取数(Gather)、x_rotate * sinComputeMul)以及InterleavedInversion(对奇数位乘以 -1,mask 为0x5555555555555555)和最终的Add求和,恰好对应公式中"拆奇偶、取负拼接、乘 cos/sin 相加"的完整计算链。

函数原型与参数说明

函数原型

cann_ops_transformer.inplace_partial_rotary_mul(x, r1, r2, *, rotary_mode="interleave", partial_slice=None) -> None

该接口位于cann_ops_transformer.ops包中(实际源码在 torch_extension/inplace_partial_rotary_mul.py,并由 torch_extension/__init__.py 导出)。接口的算子 schema 为:

inplace_partial_rotary_mul(Tensor(a!) x, Tensor r1, Tensor r2, *, str rotary_mode="interleave", int[2] partial_slice=[0, 0]) -> ()

其中Tensor(a!)表示x为可写别名(inplace 修改)语义。

参数说明

参数名参数类型可选/必选描述数据类型维度(shape)
xTensor必选待执行旋转位置编码的张量,对应公式中的 x。Inplace 模式下,计算结果直接写回该 Tensor。bfloat16、float16、float32(B, S, N, D)
r1Tensor必选位置编码张量,对应公式中的 cos 分量。bfloat16、float16、float324 维,需与 x 满足广播关系
r2Tensor必选位置编码张量,对应公式中的 sin 分量,r2 和 r1 的数据类型必须一致。bfloat16、float16、float324 维,需与 x 满足广播关系
rotary_modestr可选旋转模式。当前仅支持 "interleave",默认值为 "interleave"。--
partial_sliceList[int]可选部分旋转的切片范围 [start, end),作用于 x 的最后一维 D 维。默认值为 None,接口内部按 [0, 0] 处理。--

关于数据类型与格式,算子注册代码 inplace_partial_rotary_mul_def.cpp 明确了x支持DT_FLOAT16 / DT_FLOAT / DT_BF16cos(r1)与sin(r2)支持DT_FLOAT16 / DT_FLOAT / DT_BF16,格式统一为 ND(即内存连续布局),并带有AutoContiguous声明——这解释了文档中"不支持非连续 Tensor"的约束来源。

返回值说明

该接口无返回值(None)。计算结果直接 inplace 写回输入张量xx在计算后 shape 和数据类型保持不变,partial_slice指定范围以外的数据保持原值。这一行为在 inplace_partial_rotary_mul_infershape.cpp 中有直接体现:InferShape 将输出 shape 直接赋值为输入x的 shape(*yShape = *xShape),InferDataType 同样透传输入数据类型,确认输出即输入、不产生新张量。

自动微分说明:当x.requires_grad为 True 时,对 loss 执行.backward()即自动触发inplace_partial_rotary_mul_backward,无需手动调用反向算子。r1(cos)、r2(sin)的梯度不计算,始终为 None。反向算子inplace_partial_rotary_mul_backward支持空 Tensor 和切片长度为零的场景(执行 no-op),自动微分可正常使用。

自动微分链路的实现位于 inplace_partial_rotary_mul.py:InplacePartialRotaryMulFn继承torch.autograd.Function,forward 中通过ctx.mark_dirty(x)标记 inplace 修改并保存r1r2用于反向;backward 中调用torch.ops.cann_ops_transformer.inplace_partial_rotary_mul_backward就地计算grad_input,且返回值中除grad_input外其余均为None——这正是"r1、r2 不计算梯度"的代码级证据。此外,Python 入口 inplace_partial_rotary_mul.py 会按x.requires_grad自动分派:需要梯度时走 autograd Function 包装,否则直接调用底层算子,避免不必要的计算图开销。

约束说明

  • 该接口支持训练、推理场景下使用。
  • 该接口支持单算子模式和图模式调用。
  • 不支持非连续 Tensor。
  • x最后一维 D 大小不超过 1024,且 D 必须为 2 的倍数。该上限在 tiling 代码 inplace_partial_rotary_mul_tiling.cpp 中以D_LIMIT = 1024常量直接体现。
  • partial_slice必须包含两个整数,满足start >= 0end >= 0end <= Dend - start >= 0
  • partial_slice切片长度(即end - start)必须为 2 的倍数。当endstart相等时(切片长度为零),正向和反向计算均执行 no-op,直接返回。tiling 代码中以TILING_KEY_NOOP = 1sliceLength = 0, no computation)标记该分支(见 inplace_partial_rotary_mul_tiling.cpp)。
  • r1r2最后一维大小必须相同,且必须等于partial_slice的切片长度(即end - start)。
  • r1r2的 shape 必须与x[..., start:end]满足广播关系,且存在如下约束:
    • Ascend 950PR / Ascend 950DTr1r2的 shape 当前只支持 BSND、B1ND、B11D、111D 排布。
    • Atlas A3 训练系列产品 / Atlas A3 推理系列产品、Atlas A2 训练系列产品 / Atlas A2 推理系列产品r1r2的 shape 当前只支持 BS1D、B11D 排布。
  • x的各维度值必须大于 0;当partial_slice不是空切片时,r1r2参与计算的维度值必须大于 0。
  • 自动微分约束:仅计算x的梯度;r1r2的梯度不计算,始终为 None。因算子为输入输出同地址操作,x不能是requires_grad=True的叶子张量。

关于最后一条约束,从代码实现看,autograd 包装在 backward 阶段直接对grad_output执行就地反向(inplace_partial_rotary_mul.py),因此若x本身是requires_grad=True的叶子张量,inplace 修改会被 PyTorch 的版本计数器(version counter)机制判定为非法改写,这也是文档要求"训练时对x先做一次拷贝(如y = x * 1.0)再传入"的原因。

确定性计算

默认支持确定性计算。

调用说明

单算子模式调用

单算子模式下,接口直接完成 RoPE 编码,无需手动管理计算图:

import torch import torch_npu from cann_ops_transformer.ops import inplace_partial_rotary_mul torch_npu.npu.set_device(0) B = 2 S = 32 N = 8 D = 128 slice_start = 0 slice_end = 64 x = torch.randn(B, S, N, D, device="npu", dtype=torch.float16) r1 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) r2 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) inplace_partial_rotary_mul( x, r1, r2, rotary_mode="interleave", partial_slice=[slice_start, slice_end], )

示例要点说明:

  • x为 BSND 布局的(2, 32, 8, 128),D=128;partial_slice=[0, 64]表示仅旋转最后一维的前 64 个元素,后 64 个元素保持原值。
  • r1r2取 BS1D 排布(2, 32, 1, 64),其最后一维 64 恰等于切片长度end - start,满足上文约束;同时1维与x的 N 维满足广播关系(Atlas A2/A3 系列支持的 BS1D、B11D 排布之一)。
  • 调用后无需接收返回值,x本身即被修改;可在调用后直接打印x[0, 0, 0, :8]对比partial_slice内外的值验证效果。

训练模式调用(自动微分)

训练场景下,接口通过 autograd 包装自动衔接反向计算。注意由于 inplace 特性,需先对叶子张量x做一次拷贝:

import torch import torch_npu from cann_ops_transformer.ops import inplace_partial_rotary_mul torch_npu.npu.set_device(0) B, S, N, D = 2, 32, 8, 128 slice_start, slice_end = 0, 64 x = torch.randn(B, S, N, D, device="npu", dtype=torch.float16, requires_grad=True) r1 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) r2 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) y = x * 1.0 y.retain_grad() # 正向:自动追踪计算图(y被inplace修改,无需接收返回值) inplace_partial_rotary_mul( y, r1, r2, rotary_mode="interleave", partial_slice=[slice_start, slice_end], ) # 继续前向计算 loss = y.sum() loss.backward() # 自动调用inplace_partial_rotary_mul_backward print(y.grad.shape) print(x.grad.shape) # r1.grad, r2.grad始终为None(cos/sin不计算梯度)

此处的关键点在于:y = x * 1.0构造出非叶子的中间张量作为 inplace 修改对象,规避"x不能是requires_grad=True的叶子张量"的约束;loss.backward()会经由InplacePartialRotaryMulFn.backward自动调用反向算子inplace_partial_rotary_mul_backward(源码见 inplace_partial_rotary_mul.py),并在反向中对r1r2返回None梯度。

图模式调用

图模式(torch.compile+torchairNPU 后端)下,将算子封装进nn.Module后编译执行:

import torch import torch_npu import torchair from cann_ops_transformer.ops import inplace_partial_rotary_mul torch_npu.npu.set_device(0) B = 2 S = 32 N = 8 D = 128 slice_start = 0 slice_end = 64 class InplacePartialRotaryMulModel(torch.nn.Module): def forward(self, x, r1, r2): inplace_partial_rotary_mul( x, r1, r2, rotary_mode="interleave", partial_slice=[slice_start, slice_end], ) return x model = InplacePartialRotaryMulModel().npu() npu_backend = torchair.get_npu_backend() model = torch.compile(model, backend=npu_backend, dynamic=False) x = torch.randn(B, S, N, D, device="npu", dtype=torch.float16) r1 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) r2 = torch.randn(B, S, 1, slice_end - slice_start, device="npu", dtype=torch.float16) output = model(x, r1, r2)

从实现角度看,图模式之所以可行,得益于 torch 扩展侧对算子做了两层准备:一是注册了Meta实现(inplace_partial_rotary_mul.py),对x.dim() != 4rotary_mode != "interleave"partial_slice长度不等于 2 等非法输入在编译期即抛出明确错误;二是通过_ensure_initialized()load()预加载算子模块,避免 dynamo 跟踪期才触发torch.utils.cpp_extension.load()(见 inplace_partial_rotary_mul.py)。同时 graph_convert_inplace_partial_rotary_mul.py 负责在构图阶段完成算子到图 IR 的转换。

底层实现与调用链

Python 到 AICore 的完整调用链

该算子的完整调用链可以概括为四层:

  1. Python 接口层:inplace_partial_rotary_mul.py 负责默认参数归一化(partial_slice=None -> [0, 0])、autograd 分派,最终下发到torch.ops.cann_ops_transformer.inplace_partial_rotary_mul
  2. C++ torch 扩展层:csrc/inplace_partial_rotary_mul.cpp 校验x.dim() == 4rotary_mode,将字符串模式映射为整数(mode_map"interleave" -> 1),再调用aclnnInplacePartialRotaryMul
  3. ACLNN 计算接口层:op_api/aclnn_inplace_partial_rotary_mul.h 提供标准的GetWorkspaceSize+Execute两段式接口,rotary_modeint64_t传入(当前仅支持 interleave,即 rotary_mode=1);
  4. Host/Kernel 层:inplace_partial_rotary_mul_def.cpp 完成算子注册与 AICore 配置,inplace_partial_rotary_mul_tiling.cpp 根据输入 shape、partial_slice、dtype 计算分核与分块策略,最终由 op_kernel 下的 Kernel 代码在 NPU AICore 上执行旋转计算。

Tiling 与 Kernel 的工程细节

从 tiling 源码可以看出该算子在工程实现上的几个关键决策:

  • 多种旋转布局支持:Kernel 头文件 inplace_rotate_half.h 中定义了LAYOUT_BNSD / LAYOUT_BSND / LAYOUT_SBND / LAYOUT_R_B1SD / LAYOUT_BND / LAYOUT_NO_BROADCAST等多种布局路径,分别对应r1/r2广播形态不同时的访存策略(如 B1ND 广播需要按 batch 复用旋转系数,对应RB1sdProcess);
  • 混合精度计算:非 FP32 输入会在 Kernel 内先Cast到 FP32 参与乘加,最终以CAST_RINT模式写回原精度(见 inplace_partial_rotary_mul.h),对应 tiling 中的TILING_KEY_BFLOAT16_FLOAT32_MIXED等混合精度 tiling 分支;
  • no-op 短路:切片长度为零时 tiling 直接标记TILING_KEY_NOOP,正向与反向均跳过实际计算;
  • 多核并行分片:tiling 依据 batch、seq、head 维度做 Split 分核(TilingSplitN / TilingSplitB / TilingSplitS),每个核处理独立的 batch/head 分片,平衡负载。

示例与测试验证

仓库为该算子提供了完整的示例与单测:

  • C++ 单算子示例:examples/test_aclnn_inplace_partial_rotary_mul.cpp 与 examples/test_geir_inplace_partial_rotary_mul.cpp,分别演示 ACLNN 接口与 GEIR 图接口的直接调用;
  • Host 侧单测:tests/ut/op_host/test_inplace_partial_rotary_mul_a3_tiling.cpp 与 test_inplace_partial_rotary_mul_a5_tiling.cpp,覆盖 Atlas A3 与 A5 平台的 tiling 策略;
  • Kernel 侧单测:tests/ut/op_kernel/test_inplace_partial_rotary_mul_kernel.cpp,配套 tests/ut/op_kernel/inplace_partial_rotary_mul_data 下的gen_data.py(构造输入)、gen_golden.py(生成 golden 结果)、gen_tiling.py(生成 tiling 参数)三个脚本,构成完整的数据生成—计算—校验闭环。

总结

inplace_partial_rotary_mul通过"输入输出同地址"的 inplace 设计,配合partial_slice局部切片能力,为 transformer 类大模型在 NPU 上的旋转位置编码提供了一种低显存开销、支持自动微分的实现方案。其核心价值可以归纳为三点:

  1. 省显存、免拷贝:直接改写输入张量x,不产生额外输出,适合长序列、大 batch 场景;
  2. 局部旋转partial_slice=[start, end)精确控制旋转范围,支持"仅旋转头维度前 D/2"等常见 RoPE 变体,start == end时自动 no-op;
  3. 全链路支持:单算子、训练(autograd 自动衔接反向)、图模式(torch.compile + torchair)三种调用方式覆盖推理与训练全场景,且r1/r2作为常量参与计算、不产生梯度,符合 RoPE 的标准用法。

在接入该算子时,务必重点核对x的 BSND 布局与 D 维约束、r1/r2最后一维等于切片长度且满足目标产品的广播排布(A2/A3 系列为 BS1D、B11D;950 系列为 BSND、B1ND、B11D、111D)、以及训练场景下对叶子张量的拷贝处理,即可稳定获得正确的部分旋转位置编码结果。

  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询