CANN ops-transformer aclnnGatherPaKvCache 算子详解:PagedAttention 非连续 KV Cache 的 Gather 拼接与两段式调用实战
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
aclnnGatherPaKvCache 是 CANN ops-transformer 中为 PagedAttention(分页注意力)推理场景设计的 KV Cache 重组算子,它根据 blockTables 中的 blockId 与 seqLens 中的序列长度,把散落在多个物理块中、内存不连续的 token,从 keyCache/valueCache 搬运并拼接成一段段连续的 key/value 序列(keyRef/valueRef)。本文以 aclnnGatherPaKvCache 官方接口文档 为主体骨架,结合 GatherPaKvCache 算子模块 的算子定义、InferShape、Tiling 与 Kernel 源码展开讲解。读完本文,你将掌握该算子的完整参数语义、isSeqLensCumsum两种 seqLens 语义下的 shape 推导规则、两段式 aclnn 接口的调用流程,以及一份可直接编译运行的 C++ 样例。
算子功能与计算逻辑
功能定位:内存不连续 → 内存连续
在 PagedAttention 类大模型推理框架中,KV Cache 以"物理块"(block)为单位按需分配。一个序列的 token 往往散布在不连续的多个物理块中,无法直接作为连续的 key/value 输入喂给 Attention 内核。aclnnGatherPaKvCache 解决的就是这一"gather(收集)"环节:
根据
blockTables中的 blockId 值、seqLens中 key/value 的 seqLen,从keyCache/valueCache中将内存不连续的 token 搬运、拼接成连续的 key/value 序列。
算子模块的 README 也给出了同样的定位说明,见 attention/gather_pa_kv_cache/README.md。
计算逻辑与输出第一维的确定规则
keyRef/valueRef 的第一个维度(token 总数)取决于 seqLens 的内容,具体由属性isSeqLensCumsum决定:
- 当
isSeqLensCumsum为true(seqLens 是累加和):keyRef[dim0] = seqLens[-1],即 seqLens 的最后一个值就是输出序列总长度; - 当
isSeqLensCumsum为false(seqLens 是各 batch 的真实序列长度):keyRef[dim0] = sum(seqLens),即所有序列长度累加。
该规则与 gather_pa_kv_cache_infershape.cpp 中CheckCommonPagedCacheLoad的校验逻辑相互印证:is_seq_lens_cumsum为 true 时要求seqLens.shape[0] == blockTables.shape[0] + 1(多出一个前缀 0),为 false 时要求seqLens.shape[0] == blockTables.shape[0]。
单 token 大小限制(148k 约束)
关于 keyRef、valueRef 有一个重要的限制条件:
- 每个 token 大小控制在148k 以内。例如,对于 fp16/bf16 类型,
num_heads * head_size(keyRef/valueRef)取 128 × 576。
这一约束与内核实现直接对应:在 gather_pa_kv_cache_nd.h 中,ND 内核为搬运申请了UB_BUF_SIZE = 192 * 1024(192KB)的片上 Unified Buffer,并注释"ensure the size of one token less than 148kb",即单 token 数据必须能在 UB 缓冲内完成中转搬运。
典型 Shape 示例
官方文档给出的示例 shape 如下(batch=16,每序列最多 12 个物理块,cache 块大小 128,16 个 head,k 的 head_size=144、v 的 head_size=128):
keyCache_shape: [128, 128, 16, 144] # [num_blocks, block_size, num_heads, head_size_k] valueCache_shape: [128, 128, 16, 128] # [num_blocks, block_size, num_heads, head_size_v] blockTables_shape: [16, 12] # [batch, block_indices] seqLens_shape: [16] # [batch] keyRef_shape: [8931, 16, 144] # [num_tokens, num_heads, head_size_k] valueRef_shape: [8931, 16, 128] # [num_tokens, num_heads, head_size_v] seqOffset_shape: [16] # [batch] out1_shape: [8931, 16, 144] out2_shape: [8931, 16, 128]其中keyRef[dim0] = sum(seqLens) = 8931即对应isSeqLensCumsum = false的语义。
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品 | 不支持 |
| Atlas 训练系列产品 | 不支持 |
从算子注册看,gather_pa_kv_cache_def.cpp 中分别通过GatherPaKvCacheDefFor910b()(注册ascend910b、ascend910_93两个配置)和GatherPaKvCacheDefFor950()(注册ascend950、ascend350两个配置)完成产品适配,与上表支持范围一致。
两段式接口与函数原型
与其他 CANN aclnn 算子一致,aclnnGatherPaKvCache 采用两段式接口设计(即docs/zh/context/two_phase_api.md):必须先调用aclnnGatherPaKvCacheGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器,再调用aclnnGatherPaKvCache执行计算。
第一段接口原型:
aclnnStatus aclnnGatherPaKvCacheGetWorkspaceSize( const aclTensor *keyCache, const aclTensor *valueCache, const aclTensor *blockTables, const aclTensor *seqLens, aclTensor *keyRef, aclTensor *valueRef, const aclTensor *seqOffsetOptional, char* cacheMode, bool isSeqLensCumsum, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口原型:
aclnnStatus aclnnGatherPaKvCache( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)接口的完整实现位于 op_host/op_api/aclnn_gather_pa_kv_cache.cpp,第二段接口通过CommonOpExecutorRun(workspace, workspaceSize, executor, stream)完成 AICore 任务下发。
aclnnGatherPaKvCacheGetWorkspaceSize 参数详解
下表完整列出第一段接口的全部参数(来自官方文档,并补充了格式细节):
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| keyCache(aclTensor*) | 输入 | 表示在当前层存储的 key 向量缓存 | 当 cacheMode 为 "Norm" 时,shape 为[num_blocks, block_size, num_heads, head_size_k],数据格式必须是 ND;当 cacheMode 为 "PA_NZ" 时,shape 为[num_blocks, num_heads * head_size_k // elenum_aligned, block_size, elenum_aligned],数据格式必须是 FRACTAL_NZ | INT8、FLOAT16、BFLOAT16、FLOAT、UINT8、INT16、UINT16、INT32、UINT32、HIFLOAT8、FLOAT8_E5M2、FLOAT8_E4M3FN | ND、FRACTAL_NZ | 4 | √ |
| valueCache(aclTensor*) | 输入 | 表示在当前层存储的 value 向量缓存 | 同 keyCache 的格式要求(head_size 换成 head_size_v) | 同 keyCache | ND、FRACTAL_NZ | 4 | √ |
| blockTables(aclTensor*) | 输入 | 表示每个序列对应的物理块索引 | shape 为[batch, block_indices],其中 batch、block_indices 均须大于 0。元素取值范围为[0, num_blocks),即 blockId 可取 0 到 num_blocks-1 | INT32、INT64 | ND | 2 | × |
| seqLens(aclTensor*) | 输入 | 表示每个 batch 对应的序列长度 | shape 为[batch]或[batch + 1]。当 isSeqLensCumsum 为 false 时 shape 为[batch];为 true 时 shape 为[batch + 1]。元素取值范围为[0, num_blocks) | 与 blockTables 保持一致 | ND | 1 | × |
| keyRef(aclTensor*) | 输入/输出 | 表示 key 向量 | 当 cacheMode 为 "Norm" 时 shape 为[num_tokens, num_heads, head_size_k];为 "PA_NZ" 时 shape 为[num_tokens, num_heads * head_size_k] | 与 keyCache 保持一致 | ND | 2-3 | √ |
| valueRef(aclTensor*) | 输入/输出 | 表示 value 向量 | 当 cacheMode 为 "Norm" 时 shape 为[num_tokens, num_heads, head_size_v];为 "PA_NZ" 时 shape 为[num_tokens, num_heads * head_size_v] | 与 valueCache 保持一致 | ND | 2-3 | √ |
| seqOffsetOptional(aclTensor*) | 输入 | 如果传入,表示在从 blockTables 获取 blockId 时存在首偏移(偏移量为seqOffsetOptional[i] / block_size,i表示某一个 batch);不传入表示不需要偏移 | shape 为[batch] | 与 blockTables 保持一致 | ND | 1 | × |
| cacheMode(char*) | 输入 | 支持 ["Norm", "PA_NZ"] 两种模式,分别表示输入 keyCache 和 valueCache 数据格式是 ND、FRACTAL_NZ | - | - | - | - | - |
| isSeqLensCumsum(bool) | 输入 | 表示 seqLens 是否为累加和 | false 表示非累加和,例如 seqLens 为[1, 3, 5, 3, 7];true 表示累加和,例如 seqLens 为[0, 1, 4, 9, 12, 19],此时第 0 个元素必定是 0。累加和的seqlens[i + 1] - seqlens[i]等于非累加的seqlens[i] | bool | - | - | - |
| workspaceSize(uint64_t*) | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor(aclOpExecutor**) | 输出 | 返回 op 执行器,包含了算子计算流程 | - | - | - | - | - |
属性默认值与算子定义佐证
在 gather_pa_kv_cache_def.cpp 中可以看到两个属性均有默认值:
Attr("cache_mode").AttrType(OPTIONAL).String("Norm"):cache_mode 默认取 "Norm",即默认按 ND 格式的 KV Cache 处理;Attr("is_seq_lens_cumsum").AttrType(OPTIONAL).Bool(true):is_seq_lens_cumsum 默认取 true。
Tiling 阶段(gather_pa_kv_cache_tiling_arch35.cpp)还会根据这两个属性 + 是否传入 seqOffset 组合出 8 种 tiling key(1111/1110/1101/1100 对应 ND 四种组合,1011/1010/1001/1000 对应 PA_NZ 四种组合),见gather_pa_kv_cache_tiling.h中的tilingKeyTable,实现"同一算子、多内核形态"的按需编译。
cacheMode 对 shape 与 tiling 的影响
- Norm 模式(ND):keyCache shape 为
[num_blocks, block_size, num_heads, head_size_k],Tiling 中blockSize取自存储 shape 的第 1 维,tokenSizeK = num_heads * head_size_k(见 gather_pa_kv_cache_tiling.cpp 的CommonGatherPaKvCacheTiling); - PA_NZ 模式(FRACTAL_NZ):keyCache shape 为
[num_blocks, num_heads * head_size_k // elenum_aligned, block_size, elenum_aligned],其中elenum_aligned与元素位宽相关——b8 场景(每个数据元素位宽 8bit,如 INT8)取 32,b16 场景(如 INT16)取 16,b32 场景(如 INT32)取 8。Tiling 中blockSize取自第 2 维,并校验shape[3] * 元素字节数 == 32B(32 字节对齐),见GetInputKeyCache/GetInputValueCache。
平台差异与数据类型限制
在支持的产品上,不同架构对数据类型还有进一步限制(官方文档):
- Ascend 950PR / Ascend 950DT:允许 keyCache 为 FLOAT8_E4M3FN、valueCache 为 FLOAT16 或 BFLOAT16 的组合。
- Atlas A2 训练/推理系列产品、Atlas A3 训练/推理系列产品:
- 输入 keyCache、valueCache、keyRef、valueRef 不支持 FLOAT、UINT8、INT16、UINT16、INT32、UINT32、HIFLOAT8、FLOAT8_E5M2、FLOAT8_E4M3FN 数据类型;
- 输入 blockTables、seqLens、seqOffsetOptional 不支持 INT64 数据类型。
这与 gather_pa_kv_cache_def.cpp 的注册一致:910b 系列(对应 Atlas A2/A3)的 key/value 仅注册 FLOAT16、BFLOAT16、INT8 三种数据类型,索引类仅 INT32;而 950/350 系列注册了更完整的数据类型集合。接口层(aclnn_gather_pa_kv_cache.cpp)还会按当前 NPU 架构(DAV_3510与否)选择DTYPE_SUPPORT_LIST_A5或DTYPE_SUPPORT_LIST_910B做数据校验,并额外校验 keyCache 与 keyRef、valueCache 与 valueRef、blockTables 与 seqLens(以及 seqOffsetOptional)之间的类型一致性。
返回值与错误码
两段接口的返回值均为aclnnStatus,具体状态码可参考 aclnn 返回码(即原文档链接../../../docs/zh/context/aclnn_return_code.md)。
第一段接口完成入参校验,出现以下场景时报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 输入是空指针。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 输入数据类型不在支持的范围内。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 输入的维数不匹配。 |
对应的校验逻辑在aclnnGatherPaKvCacheGetWorkspaceSize实现中依次执行:CheckNullptr(空指针检查)→ 空 shape 提前返回 →CheckShape(keyCache/valueCache 必须 4 维、blockTables 必须 2 维、seqLens 必须 1 维、keyRef/valueRef 必须 2 或 3 维)→CheckDtypeValid(数据类型及一致性检查)。
aclnnGatherPaKvCache(第二段接口)参数说明
| 参数名 | 输入/输出 | 描述 |
|---|---|---|
| workspace | 输入 | 在 Device 侧申请的 workspace 内存地址。 |
| workspaceSize | 输入 | 在 Device 侧申请的 workspace 大小,由第一段接口 aclnnGatherPaKvCacheGetWorkspaceSize 获取。 |
| executor | 输入 | op 执行器,包含了算子计算流程。 |
| stream | 输入 | 指定执行任务的 Stream。 |
约束说明
- 确定性计算:aclnnGatherPaKvCache 默认确定性实现(确定性计算的通用约定可参考 docs/zh/context/determinism_compute.md)。
源码级实现剖析
Host 侧:InferShape 与数据流
gather_pa_kv_cache_infershape.cpp 负责 shape 推导:
- 输出 key/value 的 shape 直接继承输入 key/value(第 4、5 个输入)的 shape;
- 输入存在 UnknownRank/UnknownShape(动态场景)时输出相应置为未知;
- 校验 blockTables 维数为 2、seqLens 维数为 1,并按
is_seq_lens_cumsum校验seqLens.shape[0]与blockTables.shape[0]的关系(true:batch+1;false:batch); - 按 cache_mode 分别走
InferShape4GatherPaKvCacheNd(key/value 输出为 3 维)或InferShape4GatherPaKvCacheNz(key/value 输出为 2 维)。
Host 侧:Tiling 与分核
gather_pa_kv_cache_tiling.cpp 与 gather_pa_kv_cache_tiling_arch35.cpp 完成计算切分:
- 通用路径固定申请 16MB workspace(
ASCENDC_TOOLS_WORKSPACE); - blockDim(分核数):Norm 模式使用全部 AIV 核数;PA_NZ 模式取
min(总token数, AIV核数); - arch35 路径针对非连续 tensor(view)场景做了专门设计:通过
GetTensorInfo获取逻辑 shape 与 stride,判断 cache/ref 各轴是否连续,用nonContiguousFlag位图标记非连续形态(bit0~bit9 分别对应 keyCache、valueCache、key 输出、value 输出及各轴的非连续状态),并将各 stride 写入 tiling 数据;ND 非连续场景改为"全局 block 轮询"分核,解决 batch 数很少时按 batch 分核只用少量核的负载均衡问题。
Kernel 侧:AIV 搬运内核
gather_pa_kv_cache.cpp 是内核入口,声明为KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY),按编译期 dtype 与 tiling key 分发到:
- ND 内核
GatherPaKvCacheNd<T>(gather_pa_kv_cache_nd.h),tiling key 618(INT8)/619(FLOAT16、BF16); - NZ 内核
GatherPaKvCacheNz<T>(gather_pa_kv_cache_nz.h),tiling key 577。
ND 内核中,GetBlockCacheOffset根据cacheBlockStride(view 场景下由 host 传入的 stride)计算物理块首地址偏移,随后按blockSize × tokenSize分 cacheline 切块搬运到 192KB UB 缓冲再写回连续输出,这正是"非连续 token 搬运拼接"的落地实现。
接口层:非连续 Tensor 的零拷贝与兜底
aclnn_gather_pa_kv_cache.cpp 在入参校验后,会根据各 tensor 的 view shape/stride 判断非连续形态并分流:
- ProcessNonContiguous:cache 首轴非连续(dim0)或 ND 模式内部轴(slot/head)非连续、ref 内部轴非连续但尾轴连续等 kernel 可零拷贝处理的形态,通过
CreateView保留 view stride 直接下发给 kernel,kernel 按 stride 搬运/散写,避免物理拷贝; - ProcessContiguous:对 kernel 无法处理的形态(如 NZ 非 dim0 轴非连续、ND 尾轴非连续)先通过
l0op::Contiguous物理连续化,计算后经ViewCopy回写到原始 ref 内存,保证语义正确。
这也解释了文档参数表中 keyCache/valueCache/keyRef/valueRef 均标注"非连续 Tensor √"——该算子对非连续输入/输出有完整支持。
调用示例(可直接编译运行的 C++ 样例)
下面示例取自官方文档(与仓库 examples/test_aclnn_gather_pa_kv_cache.cpp 等价),演示了从环境初始化、构造输入输出、两段式调用到结果回拷与资源释放的完整流程:
#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_gather_pa_kv_cache.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 shapeSize = 1; for (auto i : shape) { shapeSize *= i; } return shapeSize; } 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); // 调用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); // 计算连续tensor的strides 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]; } // 调用aclCreateTensor接口创建aclTensor *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对外接口列表 // 根据自己的实际device填写deviceId int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出,需要根据API的接口自定义构造 std::vector<int64_t> keyCacheShape = {2, 2, 32, 2}; std::vector<int64_t> valueCacheShape = {2, 2, 32, 4}; std::vector<int64_t> blockTablesShape = {4,6}; std::vector<int64_t> seqLensShape = {4}; std::vector<int64_t> keyShape = {12, 32, 2}; std::vector<int64_t> valueShape = {12, 32, 4}; std::vector<int64_t> seqOffsetShape = {4}; void* keyCacheDeviceAddr = nullptr; void* valueCacheDeviceAddr = nullptr; void* blockTablesDeviceAddr = nullptr; void* seqLensDeviceAddr = nullptr; void* keyDeviceAddr = nullptr; void* valueDeviceAddr = nullptr; void* seqOffsetAddr = nullptr; aclTensor* keyCache= nullptr; aclTensor* valueCache = nullptr; aclTensor* blockTables= nullptr; aclTensor* seqLens= nullptr; aclTensor* key = nullptr; aclTensor* value= nullptr; aclTensor* seqOffset= nullptr; std::vector<uint16_t> keyCacheHostData(256, 1); std::vector<uint16_t> valueCacheHostData(512, 1); std::vector<int32_t> blockTablesHostData(24, 1); std::vector<int32_t> seqLensHostData(4, 3); std::vector<uint16_t> keyHostData(768, 0); std::vector<uint16_t> valueHostData(1536, 0); std::vector<int32_t> seqOffsetHostData(4, 2); char cacheMode[] = "Norm"; const bool isSeqLensCumsum = false; // 创建GatherPaKvCache的输入输出aclTensor ret = CreateAclTensor(keyCacheHostData, keyCacheShape, &keyCacheDeviceAddr, aclDataType::ACL_FLOAT16, &keyCache); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(valueCacheHostData, valueCacheShape, &valueCacheDeviceAddr, aclDataType::ACL_FLOAT16, &valueCache); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(blockTablesHostData, blockTablesShape, &blockTablesDeviceAddr, aclDataType::ACL_INT32, &blockTables); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(seqLensHostData, seqLensShape, &seqLensDeviceAddr, aclDataType::ACL_INT32, &seqLens); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(keyHostData, keyShape, &keyDeviceAddr, aclDataType::ACL_FLOAT16, &key); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(valueHostData, valueShape, &valueDeviceAddr, aclDataType::ACL_FLOAT16, &value); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(seqOffsetHostData, seqOffsetShape, &seqOffsetAddr, aclDataType::ACL_INT32, &seqOffset); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用CANN算子库API,需要修改为具体的Api名称 uint64_t workspaceSize = 0; aclOpExecutor* executor; // 调用aclnnGatherPaKvCache第一段接口 ret = aclnnGatherPaKvCacheGetWorkspaceSize(keyCache, valueCache, blockTables, seqLens, key , value, seqOffset, cacheMode, isSeqLensCumsum, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnGatherPaKvCacheGetWorkspaceSize 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); } // 调用aclnnGatherPaKvCache第二段接口 ret = aclnnGatherPaKvCache(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnGatherPaKvCache 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的接口定义修改 auto size = GetShapeSize(keyShape); std::vector<uint16_t> resultData(size, 0); ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), keyDeviceAddr, 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 ret); for (int64_t i = 0; i < size; i++) { LOG_PRINT("result[%ld] is: %d\n", i, resultData[i]); } // 6. 释放aclTensor和aclIntArray,需要根据具体API的接口定义修改 aclDestroyTensor(keyCache); aclDestroyTensor(valueCache); aclDestroyTensor(blockTables); aclDestroyTensor(seqLens); aclDestroyTensor(key); aclDestroyTensor(value); aclDestroyTensor(seqOffset); // 7. 释放device资源,需要根据具体API的接口定义修改 aclrtFree(keyCacheDeviceAddr); aclrtFree(valueCacheDeviceAddr); aclrtFree(blockTablesDeviceAddr ); aclrtFree(seqLensDeviceAddr ); aclrtFree(keyDeviceAddr); aclrtFree(valueDeviceAddr); aclrtFree(seqOffsetAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例要点解读
- 示例中
keyCacheShape = {2, 2, 32, 2}对应 Norm 模式的[num_blocks=2, block_size=2, num_heads=32, head_size_k=2];seqLensHostData = {3, 3, 3, 3}(非累加和,isSeqLensCumsum = false),因此输出keyShape = {12, 32, 2}(12 = 4 * 3 = sum(seqLens)); seqOffsetHostData = {2, 2, 2, 2}表示每个 batch 从blockTables取 blockId 时存在首偏移,偏移块数 =seqOffset[i] / block_size = 2 / 2 = 1个块;- workspace 按第一段接口返回的大小(可能为 0)条件性申请;
- 完整示例还可参考仓库 examples/test_aclnn_gather_pa_kv_cache.cpp,以及算子模块下的单元测试 tests/ut(覆盖 InferShape、Tiling 与 Kernel 三层)。
编译与运行
该示例属于标准 CANN aclnn 单算子调用工程,具体编译、链接与运行方式(包括头文件路径、libascendcl/libopapi等库的链接,以及 NPU 环境变量配置)请参考仓库通用文档 编译与运行样例(即原文档链接../../../docs/zh/context/compile_and_run_sample.md)。运行前需确保:
- 宿主机已安装与目标产品匹配的 CANN Toolkit 与算子包,且设备为上文"产品支持情况"中列出的支持型号;
- 编译时链接 aclnn 相关动态库,并正确包含
aclnnop/aclnn_gather_pa_kv_cache.h头文件; - 根据实际设备填写
deviceId,并根据实际的 KV Cache 规模调整各 tensor 的 shape 与数据类型。
小结
aclnnGatherPaKvCache 是 CANN ops-transformer 面向 PagedAttention 推理的关键算子之一。本文围绕 官方接口文档 完整梳理了其功能语义(blockTables + seqLens 驱动的非连续 KV Cache 搬运拼接)、isSeqLensCumsum两种 seqLens 语义、两段式接口的全部参数与错误码、平台差异与数据类型限制,并深入 算子模块源码 印证了 InferShape、Tiling 分核、ND/NZ 内核搬运及非连续 tensor 零拷贝处理等底层实现,最后给出了一份可直接落地的完整 C++ 调用示例,可作为在 NPU 上接入该算子的直接参考。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考