CANN ops-transformer 的 aclnnInplacePartialRotaryMul 算子:Inplace 部分旋转位置编码接口详解与实现剖析
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本篇文章围绕 CANN ops-transformer 仓库中 posembedding/inplace_partial_rotary_mul 模块的 aclnn 接口文档展开,系统讲解aclnnInplacePartialRotaryMul算子的功能定位、interleave 旋转位置编码的数学原理、partial_slice局部旋转机制、两段式接口的函数原型与完整参数语义、错误码表与约束说明,并结合该模块的算子定义、InferShape、Tiling 与 Kernel 源码剖析其底层实现,最后补充 Python 接口的调用方式与自动微分用法。读完本文,你将掌握如何在 Ascend NPU 上以 Inplace 方式高效完成单路旋转位置编码(RoPE)的前向计算,并理解其内部的数据切分与无操作(no-op)分支逻辑。
算子功能概述
aclnnInplacePartialRotaryMul执行单路旋转位置编码(Rotary Position Embedding,RoPE)的 Inplace 计算:直接修改输入张量xRef,不产生新的输出张量。与常规 RoPE 算子最大的差异在于两点:
- Inplace 语义:输入张量
xRef与输出共享同一块内存,计算结果直接写回xRef,省去输出张量的显存分配与搬运开销,适合作为 LLM 推理/训练中 attention 前处理链路的一环。 - 局部旋转:通过
partial_slice参数指定[start, end)范围,仅对输入张量最后一维(Head-Dim 维)范围内的数据执行旋转位置编码,范围之外的数据保持原值。这对应业界常见的"部分维度旋转"RoPE 变体(如仅对 KV 压缩后的部分通道做旋转)。
该算子位于 CANN ops-transformer 的 posembedding/inplace_partial_rotary_mul 目录下,同一目录中还提供了 PyTorch 接口(docs/torchapi_inplace_partial_rotary_mul.md)与图模式(GEIR)调用方式,三者共享同一套算子实现。
产品支持情况
依据 aclnnInplacePartialRotaryMul.md 与 README.md,产品支持情况如下:
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品(Atlas 310P 等) | 不支持 |
| Atlas 训练系列产品(Atlas 910 等) | 不支持 |
从算子定义源码 inplace_partial_rotary_mul_def.cpp 可以看到,AICore 侧仅注册了ascend910b、ascend910_93(即 Atlas A2/A3 训练系列对应 SoC)与ascend950三个平台的配置,与文档中的产品支持矩阵完全对应。其中ascend950平台通过ExtendCfgInfo("opFile.value", "inplace_partial_rotary_mul_apt")指定了独立的 apt 内核文件,对应 op_kernel/inplace_partial_rotary_mul_apt.cpp。
计算原理与公式
算子执行的旋转位置编码为 interleave 模式(rotary_mode等于 1),其计算过程与 PyTorch 语义等价:
x1 = x[..., ::2] # 偶数下标分量 x2 = x[..., 1::2] # 奇数下标分量 x_rotate = torch.cat((-x2, x1), dim=-1) # 旋转后的交错拼接 x = x * cos + x_rotate * sin即先按最后一维的奇偶下标将张量拆成两半,对调并取负号后拼接为旋转向量,再与cos、sin位置编码张量做逐元素乘加。数学表达为:
$$x_1 = x[..., ::2]$$
$$x_2 = x[..., 1::2]$$
$$x_{rotate} = \mathrm{cat}(-x_2, x_1)$$
$$x = x \cdot \cos + x_{rotate} \cdot \sin$$
partial_slice 局部旋转机制
partialSlice作用于输入张量的最后一维(D 维),以左闭右开区间[start, end)的形式指定需要旋转的范围:
- 不传值(或 Python 接口传入
None)时,默认按[0, 0]处理,即整条 D 维参与旋转编码; start与end相等(切片长度为 0)时,不执行旋转位置编码,直接返回(no-op);- 其余位置(
[0, start)与[end, D))的数据保持原值不变。
被旋转的局部张量为x[..., start:end],cos与sin只作用于这一段,其最后一维大小必须等于切片长度end - start。
输入张量xRef采用 BSND 维度排布:B(Batch)为批量大小,S(Seq-Length)为序列长度,N(Head-Num)为多头数量,D(Head-Dim)为每个头的隐藏维度大小。partial_slice的区间即落在 D 维上。
两段式接口与函数原型
与 CANN 大多数 aclnn 算子一致,aclnnInplacePartialRotaryMul采用两段式接口(详见 docs/zh/context/two_phase_api.md):
- 先调用
aclnnInplacePartialRotaryMulGetWorkspaceSize完成入参校验与 workspace 大小计算,获取workspaceSize与executor; - 再调用
aclnnInplacePartialRotaryMul传入 workspace 与 executor 执行实际计算。
aclnnStatus aclnnInplacePartialRotaryMulGetWorkspaceSize( const aclTensor *xRef, const aclTensor *cos, const aclTensor *sin, int64_t rotary_mode, const aclIntArray *partialSlice, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnInplacePartialRotaryMul( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)从 aclnn_inplace_partial_rotary_mul.cpp 的实现可以看到,这两个 aclnn 接口是 aclnnInner 内部接口的轻量转发封装,真实的入参校验、workspace 计算与执行逻辑在算子库内部完成。
aclnnInplacePartialRotaryMulGetWorkspaceSize 参数说明
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| xRef | 输入 | 待执行旋转位置编码的张量,公式中的 x。Inplace 模式,xRef 同时作为输出写入结果 | - | BFLOAT16、FLOAT16、FLOAT32 | ND | 4 | × |
| cos | 输入 | 位置编码张量,公式中的 cos | 与 xRef 数据类型一致,或者为 FLOAT32 | BFLOAT16、FLOAT16、FLOAT32 | ND | 4 | × |
| sin | 输入 | 位置编码张量,公式中的 sin | 与 xRef 数据类型一致,或者为 FLOAT32 | BFLOAT16、FLOAT16、FLOAT32 | ND | 4 | × |
| rotary_mode | 输入 | 旋转模式 | 0 为 half 模式,1 为 interleave 模式,当前仅支持 interleave 模式 | INT64 | - | - | - |
| partialSlice | 输入 | 部分旋转的切片范围 [start, end),作用于最后一维 | 不传值则默认整 D 轴做旋转编码,start 和 end 相等时则不做旋转编码 | INT64 数组 | - | - | - |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含算子计算流程 | - | - | - | - | - |
返回值与错误码
第一段接口返回aclnnStatus状态码(详见 docs/zh/context/aclnn_return_code.md)。第一段接口完成入参校验,出现以下场景时报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入的 xRef、cos 或 sin 是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 传入的 xRef、cos、sin 的数据类型不在支持范围内,或 cos 与 sin 的数据类型不一致,或精度组合不满足要求 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 传入的 rotary_mode 不为 1(仅支持 interleave 模式) |
| ACLNN_ERR_PARAM_INVALID | 161002 | 传入的 partialSlice 长度不为 2,或取值范围不合法 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 传入的 xRef、cos、sin 的形状不满足约束(维度不为 4,或 cos 与 sin 形状不一致,或 xRef 最后一维大小超过 1024,或 xRef 最后一维不是 2 的倍数,或 partialSlice 切片长度不是 2 的倍数) |
其中"精度组合不满足要求"与算子定义源码 inplace_partial_rotary_mul_def.cpp 中声明的数据类型矩阵一致:x支持 FP16/FLOAT32/BF16;cos、sin额外兼容 FLOAT32(即当 x 为半精度时,cos/sin 可抬升为 FLOAT32 参与计算,也就是 Kernel 侧"mixed"模板的来源)。
aclnnInplacePartialRotaryMul 参数说明
| 参数名 | 输入/输出 | 描述 |
|---|---|---|
| workspace | 输入 | 在 Device 侧申请的 workspace 内存地址 |
| workspaceSize | 输入 | 在 Device 侧申请的 workspace 大小,由第一段接口 aclnnInplacePartialRotaryMulGetWorkspaceSize 获取 |
| executor | 输入 | op 执行器,包含了算子计算流程 |
| stream | 输入 | 指定执行任务的 Stream 流 |
约束说明
使用该接口时必须遵守以下约束(依据 aclnnInplacePartialRotaryMul.md):
- 确定性计算:
aclnnInplacePartialRotaryMul默认为确定性实现(确定性计算的一般说明可参考 docs/zh/context/determinism_compute.md)。 - 不支持非连续(non-contiguous)Tensor。
- 仅支持 interleave 模式(
rotary_mode = 1),half 模式(0)当前未开放。 - Inplace 执行:输入 xRef 和输出共享同一个 Tensor,计算结果直接写回输入 xRef。
- 输入张量 xRef 的 shape 为 BSND 排布,各 shape 约束如下:
- xRef 最后一维(D)大小不超过 1024;
- interleave 模式下 xRef 最后一维(D)必须为 2 的倍数,
partialSlice切片长度(partialSlice[1] - partialSlice[0])也必须是 2 的倍数; - cos、sin 最后一维大小必须相同,且必须等于
partialSlice的切片长度(partialSlice[1] - partialSlice[0]); - cos/sin 的 shape 必须与 xRef 满足 广播关系,且存在平台差异:
- Ascend 950PR / Ascend 950DT:cos/sin 的 shape 当前只支持 BSND、B1ND、B11D、111D 四种排布;
- Atlas A3 / Atlas A2 训练与推理系列产品:cos/sin 的 shape 当前只支持 BS1D、B11D 两种排布,即要求 B 轴保持相等。
- partialSlice 取值范围:
sliceStart ≥ 0,sliceEnd ≥ 0,sliceEnd ≤ xRef 最后一维(D)大小,sliceLength = sliceEnd - sliceStart >= 0;当sliceEnd与sliceStart相同时,不做旋转位置编码,直接返回。
完整调用示例(C++)
以下示例来自官方文档,可直接参考 examples/test_aclnn_inplace_partial_rotary_mul.cpp 与编译与运行样例进行编译执行。示例中x的 shape 为{96, 1, 1, 512}(B=96, S=1, N=1, D=512),cos/sin的 shape 为{96, 1, 1, 64},partialSlice = {448, 512},即只对最后一维的[448, 512)共 64 个元素做旋转编码。
#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_inplace_partial_rotary_mul.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::vector<int64_t>& shape) { int64_t shape_size = 1; for (auto i : shape) { shape_size *= i; } return shape_size; } int Init(int32_t deviceId, aclrtStream* stream) { 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 = 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::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr, aclDataType dataType, aclTensor** tensor) { auto size = GetShapeSize(shape) * sizeof(T); 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); 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); std::vector<int64_t> strides(shape.size(), 1); for (int64_t i = shape.size() - 2; i >= 0; i--) { strides[i] = shape[i + 1] * strides[i + 1]; } *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1. 固定写法,device/stream初始化,参考acl API手册 int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == 0, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入,需要根据API的接口定义构造 std::vector<int64_t> xShape = {96, 1, 1, 512}; std::vector<int64_t> cosShape = {96, 1, 1, 64}; std::vector<int64_t> sinShape = {96, 1, 1, 64}; int64_t rotary_mode = 1; std::vector<int64_t> partialSlice = {448, 512}; void* xDeviceAddr = nullptr; void* cosDeviceAddr = nullptr; void* sinDeviceAddr = nullptr; aclTensor* x = nullptr; aclTensor* cos = nullptr; aclTensor* sin = nullptr; // Create host data (example data) std::vector<float> xHostData = { /* fill with your data */ }; std::vector<float> cosHostData = { /* fill with your data */ }; std::vector<float> sinHostData = { /* fill with your data */ }; ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT, &x); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(cosHostData, cosShape, &cosDeviceAddr, aclDataType::ACL_FLOAT, &cos); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(sinHostData, sinShape, &sinDeviceAddr, aclDataType::ACL_FLOAT, &sin); CHECK_RET(ret == ACL_SUCCESS, return ret); aclIntArray* partialSliceArray = aclCreateIntArray(partialSlice.data(), partialSlice.size()); // 3. 调用CANN算子库API uint64_t workspaceSize = 0; aclOpExecutor* executor; // 调用aclnnInplacePartialRotaryMul第一段接口 ret = aclnnInplacePartialRotaryMulGetWorkspaceSize(x, cos, sin, rotary_mode, partialSliceArray, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("GetWorkspaceSize failed. ERROR: %d\n", ret); return ret); 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;); } // 调用aclnnInplacePartialRotaryMul第二段接口 ret = aclnnInplacePartialRotaryMul(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("InplacePartialRotaryMul 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); // Cleanup aclDestroyIntArray(partialSliceArray); aclDestroyTensor(x); aclDestroyTensor(cos); aclDestroyTensor(sin); aclrtFree(xDeviceAddr); aclrtFree(cosDeviceAddr); aclrtFree(sinDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }调用流程可归纳为四步:初始化 ACL 环境 → 构造 aclTensor 输入与 aclIntArray 属性 → 两段式调用算子 → 同步并释放资源。注意xHostData需按真实推理数据填充,且示例中 cos/sin 均以ACL_FLOAT构造——这正是文档允许的"cos/sin 可为 FLOAT32 精度"组合。
从源码看实现原理
算子定义:数据类型与属性默认值
inplace_partial_rotary_mul_def.cpp 通过OpDef注册算子,其中:
- 输入
x、cos、sin均声明REQUIRED,格式为FORMAT_ND,并通过AutoContiguous()声明自动连续化要求(与文档"不支持非连续 Tensor"的约束对应); - 输入
cos、sin的数据类型列表为{FP16, FP32, BF16, FP32, FP32},即在 x 为 FP16/BF16 时允许 cos/sin 以 FP32 参与(mixed 精度模板),在 x 为 FP32 时 cos/sin 必须同为 FP32; - 属性
rotary_mode类型OPTIONAL、默认值 0;属性partial_slice类型OPTIONAL、默认值{0, 0},与文档"不传值则默认整 D 轴做旋转编码"的语义一致; - 输入与输出的张量名同为
"x",直观体现了 Inplace 输入输出同地址的语义。
InferShape:输出即输入
inplace_partial_rotary_mul_infershape.cpp 的实现非常简洁:*yShape = *xShape,输出 shape 直接拷贝输入 shape;输出数据类型同样直接继承输入数据类型。这也从推理层面印证了 Inplace 语义——输出张量与输入张量 shape、dtype 完全一致,计算结果原地写回。
Tiling:数据切分与 no-op 分支
tiling 头文件 中定义了InplacePartialRopeRegbaseTilingData,包含 B、S、N、D、blockNumB/S/N、blockFactorB/S/N、ubLoopNum/ubFactor/ubTailFactor(B/S/N 三个维度的核间与核内切分因子)、sliceStart/sliceEnd/sliceLength以及 A3 平台专用的coreTUbLoopTime、ubFactor等字段,同时定义了InplacePartialRotaryPosEmbeddingMode枚举(HALF=0、INTERLEAVE=1、QUARTER=2、DEEPSEEK_INTERLEAVE=3)与InplacePartialRopeLayout布局枚举(NO_BROADCAST、BROADCAST_BSN、BSND、SBND、BNSD)。
Tiling 主流程DoTiling()中有两个值得注意的点:
- no-op 提前返回:当切片为空(
sliceStart == sliceEnd)时,直接构造一份 sliceLength=0 的 tiling 数据、SetBlockDim(1)、设置专门的 tiling key(20040),并申请固定大小的 workspace 后返回,kernel 侧会检测sliceLength == 0直接跳过计算; - 按广播布局分发模板:根据 x 与 cos 的 shape 关系判定
layout_(如 BSND 的"aba_and_ba"、B11D 的"ab"等),再调用对应子类的SplitCore()、SplitUb()、ComputeUbFactor()完成核间/核内切分。
Kernel:按 TilingKey 分发
inplace_partial_rotary_mul.cpp 是 AscendC 编写的 AICore 内核入口,参数为x, cos, sin, y, workspace, tiling。内核首先读取 tiling 数据,若headDim == 0(即 no-op 场景)直接返回;随后按 tiling key 分发到不同的处理模板:
- TILING_KEY1/2(1/2):走
InplacePartialRotaryMulABA<DTYPE_X, true/false>模板; - 2000/2010/2020 等:按 S、BS、BSN 三种切分维度 × half/bf16/float 三种数据类型分发到
InterleavedSplitS、InterleavedSplitBS、InterleavedSplitBSN; - 2001/2011 等:对应带 pad 的变体
InterleavedSplitSPad、InterleavedSplitBSPad、InterleavedSplitBSNPad; - 2030/2040、2130/2140 等:对应 mixed 精度模板(cos/sin 为 FP32、x 为 half/bf16),文件分布在 op_kernel 下的
inplace_rotate_interleaved_split_*_mixed.h系列。
由此可以看出,内核针对不同的广播布局(S/BS/BSN 切分)、数据类型组合(同精度 / mixed 精度)以及是否 pad 设计了多套模板,由 tiling 阶段根据实际 shape 与平台信息选定,这也是该算子在不同 shape 组合下保持高效的原因。
Python 接口与自动微分
除了 aclnn C++ 接口,该算子还封装了 PyTorch 接口cann_ops_transformer.inplace_partial_rotary_mul(详见 docs/torchapi_inplace_partial_rotary_mul.md),封装实现位于 torch_extension/inplace_partial_rotary_mul.py。
cann_ops_transformer.inplace_partial_rotary_mul(x, r1, r2, *, rotary_mode="interleave", partial_slice=None) -> None参数语义与 aclnn 接口一一对应:x对应 xRef(BSND 排布,bfloat16/float16/float32),r1对应 cos,r2对应 sin,rotary_mode仅支持"interleave"(默认值即 interleave),partial_slice默认None,内部按[0, 0]处理。约束与 aclnn 完全一致(D ≤ 1024、D 为 2 的倍数、切片长度为 2 的倍数、r1/r2 最后一维等于切片长度等)。
单算子模式调用
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], )该接口无返回值(None),x在计算后 shape 与 dtype 保持不变,partial_slice指定范围以外的数据保持原值。
训练模式调用(自动微分)
当x.requires_grad=True时,接口内部会走InplacePartialRotaryMulFn(一个torch.autograd.Function):正向调用算子本体并将r1/r2保存;反向时自动调用inplace_partial_rotary_mul_backward计算x的梯度。r1(cos)、r2(sin)的梯度不计算、始终为 None——这与 RoPE 的标准用法一致(位置编码视为常量)。
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不计算梯度)从 inplace_partial_rotary_mul.py 源码可以看到,InplacePartialRotaryMulFn.forward通过ctx.mark_dirty(x)声明 inplace 修改,backward将 grad_output 连续化后调用反向算子;反向算子同样支持空 Tensor 与切片长度为零的场景(执行 no-op),因此自动微分可正常使用。需要注意:因算子为输入输出同地址操作,x不能是requires_grad=True的叶子张量。
图模式调用
Python 接口同样支持通过torch.compile+torchair后端以图模式运行(另有图模式 C++ 调用样例 examples/test_geir_inplace_partial_rotary_mul.cpp,算子 IR 定义见 op_graph/inplace_partial_rotary_mul_proto.h):
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_extension/graph_convert_inplace_partial_rotary_mul.py)负责在 torchair 图模式下将高层调用转换为底层 GEIR 算子节点。
总结
aclnnInplacePartialRotaryMul是 CANN ops-transformer 在 posembedding 场景下提供的一个"小而精"的算子:它以 Inplace 语义省去输出显存开销,以partial_slice支持仅对 Head-Dim 的局部区间做旋转位置编码,以 interleave 模式覆盖主流 RoPE 变体。接口层面提供 aclnn 两段式 C++ 接口、PyTorch 单算子/训练/图模式三种调用方式;实现层面则由 OpDef 定义、InferShape(输出即输入)、按广播布局分发的多模板 Tiling 以及按 TilingKey 分发的 AscendC Kernel 共同构成,并在空切片场景下通过 no-op 分支零开销返回。
在模型部署或训练脚本中,当你的 LLM 推理链路需要对 KV 缓存或中间激活张量的部分通道做原地旋转编码时,可以直接复用本算子的 aclnn 或 Python 接口,并严格遵循 D ≤ 1024、切片长度为 2 的倍数、cos/sin 最后一维等于切片长度等约束,即可获得与框架内建 RoPE 语义一致的确定性计算结果。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考