CANN ops-transformer 算子解析:NsaCompressWithCache 推理阶段 NSA KV 压缩算子实战指南
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本篇技术指南围绕 CANN ops-transformer 开源仓库中的NsaCompressWithCache算子展开,系统讲解其在 Native-Sparse-Attention(NSA)大模型推理场景下压缩 KV Cache 的算法原理、aclnn 两段式调用接口、全部参数与约束、错误码语义以及完整的可编译示例。读完本文,你将能够在 Atlas A2/A3 训练推理系列产品与 Kirin X90/Kirin 9030 处理器上,正确构造并调用aclnnNsaCompressWithCache接口完成推理阶段 KV 压缩,并理解其 tiling 分核与 kernel 广播计算的底层实现机制。
算子定位与产品支持情况
NsaCompressWithCache是 CANN ops-transformer 注意力(attention)子目录下的一个算子,位于 attention/nsa_compress_with_cache。它用于Native-Sparse-Attention 推理阶段的 KV 压缩:每次推理每个 batch 会产生一个新的 token,每当某个 batch 的 token 数量凑满一个compressBlockSize时,该算子会将该 batch 的后compressBlockSize个 token 压缩成一个compress_token,从而将 PagedAttention 场景下不断增长的 KV Cache 转化为稀疏注意力所需的压缩表示。
产品支持矩阵
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | × |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
| Kirin X90 处理器系列产品 | √ |
| Kirin 9030 处理器系列产品 | √ |
平台支持情况在算子注册源码 nsa_compress_with_cache_def.cpp 中也有对应体现:该算子通过AICore().AddConfig("ascend910b")、AddConfig("ascend910_93")注册 Atlas A2/A3 平台,并通过GetKirinCoreConfig()以kirinx90、kirin9030两个配置名注册 Kirin 平台。值得注意的一个平台差异是:Kirin X90/Kirin 9030 处理器系列产品不支持 BFLOAT16,其配置中 input/weight/outputCache 均仅注册了DT_FLOAT16(见 nsa_compress_with_cache_def.cpp)。
功能说明与压缩算法流程
算子功能
用于 Native-Sparse-Attention 推理阶段的 KV 压缩。推理时每个 batch 每步只产生一个新的 token,当某个 batch 的 token 数量满足压缩触发条件时,算子将该 batch 末尾的compressBlockSize个 token 与压缩权重weight相乘并求和,压缩为一个compress_token写入压缩缓存。
算法流程(四步)
- 检查序列长度:遍历
actSeqLenOptional,检查是否存在满足s >= compressBlockSize且(s - compressBlockSize) % stride == 0的序列长度(其中s为当前 batch 的 token 长度); - 定位待压缩数据:找到满足序列长度条件的
batchIdx,根据blockTableOptional找到该 batch 的后compressBlockSize个 token 作为待压缩数据; - 执行压缩算法:对待压缩 token 与压缩权重执行加权求和(乘加归约);
- 写回压缩缓存:根据
slotMapping将压缩结果写回到outputCacheRef中对应位置。
计算公式
$$ compressIdx = (s - compressBlockSize) / stride $$
$$ outputCacheRef[slotMapping[i]] = input[compressIdx \times stride : compressIdx \times stride + compressBlockSize] \times weight[:] $$
其中s是当前 batch 的 token 长度,compressBlockSize是压缩滑窗大小,stride是两次压缩滑窗的间隔。每次压缩产出一个新的compress_token,压缩结果被写入outputCacheRef中由slotMapping[i]指定的位置。
从 kernel 实现看,上述公式的底层计算在 op_kernel/nsa_compress_with_cache.cpp 中通过 AscendC 算子编程完成:KV 数据先由 FP16/BF16Cast为 FP32 以提高累加精度(ComputeSubTile),然后与广播展开的权重做逐元素Mul,再通过 ReduceBlock 将compressBlockSize个 token 的乘积结果归约到 1 个 token,最后Cast(CAST_RINT舍入)回 FP16/BF16 输出。
参数说明
算子在底层 GE 图中的输入输出与属性定义见 nsa_compress_with_cache_def.cpp。下表汇总了算子层面对外暴露的全部参数:
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| s | 属性 | 当前 batch 的 token 长度。 | INT64 | - |
| compressBlockSize | 属性 | 压缩滑窗大小。 | INT64 | - |
| stride | 属性 | 两次压缩滑窗间隔大小。 | INT64 | - |
| weight | 输入 | k/v 值的压缩 weight。 | BFLOAT16、FLOAT16 | ND |
| input | 输入 | k/v 值的 cache。 | BFLOAT16、FLOAT16 | ND |
| slotMapping | 输入 | 每个 batch 尾部压缩数据存储的位置的索引。 | INT32 | ND |
| outputCacheRef | 输入/输出 | 输出的 cache。 | BFLOAT16、FLOAT16 | ND |
需要特别说明的是:
- Kirin X90/Kirin 9030 处理器系列产品不支持 BFLOAT16,该平台下
input/weight/outputCache仅支持 FLOAT16; - 从 nsa_compress_with_cache_def.cpp 可以看到,
act_seq_len与block_table在底层注册为 OPTIONAL(可选)输入,其中act_seq_len还带有ValueDepend(OPTIONAL)标记,表示其数值会参与 tiling 决策;layout、compress_block_size、compress_stride、act_seq_len_type、page_block_size均注册为属性。
aclnn 接口参数(GetWorkspaceSize 阶段)
aclnnNsaCompressWithCache采用 CANN 标准的两段式(two-phase)接口设计,完整说明见 两段式接口文档。第一段接口aclnnNsaCompressWithCacheGetWorkspaceSize完成入参校验、形状推导与 workspace 计算,其函数原型如下:
aclnnStatus aclnnNsaCompressWithCacheGetWorkspaceSize( const aclTensor *input, const aclTensor *weight, const aclTensor *slotMapping, const aclIntArray *actSeqLenOptional, const aclTensor *blockTableOptional, char *layoutOptional, int64_t compressBlockSize, int64_t compressStride, int64_t actSeqLenType, int64_t pageBlockSize, aclTensor *outputCache, uint64_t *workspaceSize, aclOpExecutor **executor)各参数详细说明如下:
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续 Tensor |
|---|---|---|---|---|---|---|---|
| input | 输入 | 待压缩张量。 | 不支持空 Tensor;与 weight 满足 broadcast 关系,input 的第三维大小与 weight 的第二维大小相等;headDim 是 16 的整数倍且 ≤256;headNum≤64,且 headNum>50 时 headNum%2=0。N 表示多头数、D 表示隐藏层最小单元尺寸。 | FLOAT16、BFLOAT16 | ND | [blockNum, pageBlockSize, N, D]、[TND] | √ |
| weight | 输入 | 压缩的权重。 | 不支持空 Tensor;数据类型与 input 保持一致。 | FLOAT16、BFLOAT16 | ND | [compressBlockSize, N] | √ |
| slotMapping | 输入 | 每个 batch 尾部压缩数据存储的位置的索引。 | 不支持空 Tensor;值无重复,否则会导致计算结果不稳定。 | INT32 | ND | [B] | × |
| actSeqLenOptional | 可选输入 | 每个 Batch 对应的 S 大小。 | 在 TND 排布场景下需要该输入,其余场景输入 nullptr;值不应超过序列最大长度。S 表示输入样本序列长度。 | INT64 | ND | [B] | - |
| blockTableOptional | 可选输入 | PageAttention 中 KV 存储使用的 block 映射表。 | 不使用该功能可传入 nullptr;值不超过 blockNum,否则会发生越界。 | INT32 | ND | [batch, blockNumPerBatch] | - |
| layoutOptional | 可选输入 | 输入 input 的数据排布格式。 | 当前仅支持 "TND";当传入 blockTableOptional 时此参数无效,否则为必选参数。T 是 B 和 S 合轴紧密排列的数据(每个 batch 的 actSeqLen)。 | STRING | - | - | - |
| compressBlockSize | 输入 | 压缩滑窗大小。 | 必须是 16 的整数倍,且 compressBlockSize≥compressStride,compressBlockSize≤64。 | INT64 | - | - | - |
| compressStride | 输入 | 两次压缩间的滑窗间隔大小。 | 仅支持取值 16、32、48、64。 | INT64 | - | - | - |
| actSeqLenType | 输入 | actSeqLenOptional 的不同表达形式。 | actSeqLenOptional 有输入时生效,可取值 0 或 1:0 代表 actSeqLenOptional 中数值为前继 batch 序列大小的 cumsum(累积和)结果,1 代表其中数值为每个 batch 的序列大小;当前仅支持 1。 | INT64 | - | - | - |
| pageBlockSize | 输入 | page attention 场景下 page 的 blocksize 大小。 | 只能是 64 或者 128。 | INT64 | - | - | - |
| outputCache | 输出 | 压缩之后的 cache。 | 数据类型与 input 保持一致。 | FLOAT16、BFLOAT16 | ND | [result_len, N, D] | × |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小。 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含了算子计算流程。 | - | - | - | - | - |
关于输入排布需要区分两种场景:PageAttention 场景下(传入blockTableOptional),input的 shape 支持[blockNum, pageBlockSize, N, D];其余场景(TND 排布,传入actSeqLenOptional与layoutOptional="TND"),input的 shape 支持[T, N, D]。
第二段接口与返回值
aclnnStatus aclnnNsaCompressWithCache( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)第二段接口参数均为输入:workspace为在 Device 侧申请的 workspace 内存地址,workspaceSize由第一段接口获取,executor为第一段接口返回的算子执行器(包含算子计算流程),stream为指定执行任务的 Stream。
两个接口的返回值均为aclnnStatus状态码,具体含义参见 aclnn 返回码说明。第一段接口完成入参校验,以下场景会报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 计算输入和必选计算输出是空指针。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 计算输入和输出的数据类型和格式不在支持的范围内。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | input、weight、outputCache 为空 tensor。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | input 和 weight 不满足 broadcast 关系,即 input 的第三维大小与 weight 的第二维大小不相等。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | activeNum、expertNum、expertCapacity 的值小于 0。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | compress_block_size、compress_stride 不是 16 的整数倍。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | seq_lens_type != 1,或者 layout 取值不是 BSH、SBH、BSND、BNSD、TND 中的一个。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | page_block_size 取值不是 64 或者 128。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | headDim 未对齐 16。 |
从 aclnn_nsa_compress_with_cache.cpp 的 L2 接口实现可以看到校验链路:先做必要指针判空与空 tensor 检查,再校验 ND/NCL/NCHW/NCDHW 等存储格式,随后通过OP_CHECK_DTYPE_NOT_SUPPORT校验 FLOAT16/BF16 数据类型并保证 input、weight、outputCache 三者数据类型一致(InputDtypeCheck),最后对非连续输入执行l0op::Contiguous转换,再调用 L0 层l0op::NsaCompressWithCache完成算子调度。
约束说明
使用该算子必须满足以下约束:
input和weight满足 broadcast 关系,input的第三维大小与weight的第二维大小相等;compressBlockSize、stride必须是 16 的整数倍,且compressBlockSize >= stride,compressBlockSize <= 64;actSeqLenType目前仅支持取值 1;layoutOptional取值可以是 BSH、SBH、BSND、BNSD、TND,但当前不会生效;pageBlockSize只能是 64 或者 128;headDim是 16 的整数倍,且headDim <= 256;- 不支持 input/weight/outputCache 为空输入;
slotMapping的值无重复,否则会导致计算结果不稳定;blockTableOptional的值不超过 blockNum,否则会发生越界;actSeqLenOptional的值不应该超过序列最大长度;headNum <= 64,且headNum > 50时headNum % 2 = 0;- 确定性计算:
aclnnNsaCompressWithCache默认确定性实现; outputCache的 N 和 D 与 input 一致,且要满足result_len > (blockNum * pageBlockSize - compressBlockSize) / compressStride。
上述约束在 tiling 阶段同样会被校验。tiling 入口 nsa_compress_with_cache_tiling.cpp 在进入实际 tiling 逻辑前会先执行CheckParams(校验 shape 与 attr 可获取)与IsEmptyInput(校验 input/weight/outputCache 三者的 shape size 均非 0),任一不满足直接返回失败;参数取值范围类约束(16 对齐、pageBlockSize∈{64,128}、headDim≤256 等)则由 TilingPrepareForNsaCompressWithCache 读取平台信息(AIV 核数、UB 内存大小)后进入通用 tiling 实现统一校验。
调用说明与完整示例
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn 接口 | examples/test_aclnn_nsa_compress_with_cache.cpp | 通过aclnnNsaCompressWithCache接口方式调用 NsaCompressWithCache 算子,接口详细说明见 docs/aclnnNsaCompressWithCache.md。 |
完整的可编译调用示例位于 examples/test_aclnn_nsa_compress_with_cache.cpp,其编译与执行流程可参考 编译与运行样例。核心流程如下,示例以一个 PageAttention 场景(batch_size=4、headNum=24、headDim=192、pageBlockSize=128、compressBlockSize=32、compressStride=16、maxSeqLen=512)演示完整调用:
#include "acl/acl.h" #include "aclnnop/aclnn_nsa_compress_with_cache.h" #include <iostream> #include <vector> #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); // 调用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() { // 输入shape相关参数设置 constexpr int64_t compress_block_size = 32; constexpr int64_t compress_stride = 16; constexpr int64_t heads_num = 24; constexpr int64_t heads_dim = 192; constexpr int64_t batch_size = 4; constexpr int64_t page_block_size = 128; constexpr int64_t max_seq_len = 512; constexpr int64_t result_len = 512; constexpr int64_t block_num_per_batch = max_seq_len / page_block_size; constexpr int64_t blocks_num = block_num_per_batch * batch_size; // 1. 固定写法,device/stream初始化,参考acl对外接口列表 // 根据自己的实际device填写deviceId int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); // check根据自己的需要处理 CHECK_RET(ret == 0, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出,需要根据API的接口定义构造 std::vector<int64_t> inputShape = {blocks_num, page_block_size, heads_num, heads_dim}; std::vector<int64_t> weightShape = {compress_block_size, heads_num}; std::vector<int64_t> slotMappingShape = {batch_size}; std::vector<int64_t> outputCacheRefShape = {result_len, heads_num, heads_dim}; std::vector<int64_t> actSeqLenShape = {batch_size}; std::vector<int64_t> blockTableShape = {batch_size, block_num_per_batch}; void *inputDeviceAddr = nullptr; void *weightDeviceAddr = nullptr; void *slotMappingDeviceAddr = nullptr; void *outputCacheRefDeviceAddr = nullptr; void *actSeqLenDeviceAddr = nullptr; void *blockTableDeviceAddr = nullptr; aclTensor *input = nullptr; aclTensor *weight = nullptr; aclTensor *slotMapping = nullptr; aclTensor *outputCacheRef = nullptr; aclIntArray *actSeqLen = nullptr; aclTensor *blockTable = nullptr; std::vector<aclFloat16> inputHostData(inputShape[0] * inputShape[1] * inputShape[2] * inputShape[3], aclFloatToFloat16(1.0)); std::vector<aclFloat16> weightHostData(weightShape[0] * weightShape[1], aclFloatToFloat16(1.0)); std::vector<int32_t> slotMappingHostData(slotMappingShape[0], 0); std::vector<aclFloat16> outputCacheRefHostData(outputCacheRefShape[0] * outputCacheRefShape[1] * outputCacheRefShape[2], aclFloatToFloat16(1.0)); std::vector<int64_t> actSeqLenHostData(actSeqLenShape[0], 0); std::vector<int32_t> blockTableHostData(blockTableShape[0] * blockTableShape[1]); actSeqLenHostData[0]=32; // 创建self aclTensor ret = CreateAclTensor(inputHostData, inputShape, &inputDeviceAddr, aclDataType::ACL_FLOAT16, &input); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(weightHostData, weightShape, &weightDeviceAddr, aclDataType::ACL_FLOAT16, &weight); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(slotMappingHostData, slotMappingShape, &slotMappingDeviceAddr, aclDataType::ACL_INT32, &slotMapping); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(outputCacheRefHostData, outputCacheRefShape, &outputCacheRefDeviceAddr, aclDataType::ACL_FLOAT16, &outputCacheRef); CHECK_RET(ret == ACL_SUCCESS, return ret); actSeqLen = aclCreateIntArray(actSeqLenHostData.data(), actSeqLenHostData.size()); ret = CreateAclTensor(blockTableHostData, blockTableShape, &blockTableDeviceAddr, aclDataType::ACL_INT32, &blockTable); CHECK_RET(ret == ACL_SUCCESS, return ret); char layout[4] = "TND"; int64_t actSeqLenType = 1; // 3. 调用CANN算子库API,需要修改为具体的API uint64_t workspaceSize = 0; aclOpExecutor* executor; // 调用aclnnNsaCompressWithCache第一段接口 ret = aclnnNsaCompressWithCacheGetWorkspaceSize(input, weight, slotMapping, actSeqLen, blockTable, layout, compress_block_size, compress_stride, actSeqLenType, page_block_size, outputCacheRef, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnNsaCompressWithCacheGetWorkspaceSize 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;); } // 调用aclnnNsaCompressWithCache第二段接口 ret = aclnnNsaCompressWithCache(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnNsaCompressWithCache 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(outputCacheRefShape); std::vector<aclFloat16> resultData(size, 0); ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(aclFloat16), outputCacheRefDeviceAddr, size * sizeof(aclFloat16), 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 = heads_dim * heads_num - 16; i < heads_dim * heads_num + 16; i++) { printf("outputCache[%ld]:%f\n", i, aclFloat16ToFloat(resultData[i])); } // 6. 释放aclTensor,需要根据具体API的接口定义修改 aclDestroyTensor(input); aclDestroyTensor(weight); aclDestroyTensor(slotMapping); aclDestroyTensor(outputCacheRef); aclDestroyIntArray(actSeqLen); aclDestroyTensor(blockTable); // 7. 释放device资源,需要根据具体API的接口定义修改 aclrtFree(inputDeviceAddr); aclrtFree(weightDeviceAddr); aclrtFree(slotMappingDeviceAddr); aclrtFree(outputCacheRefDeviceAddr); aclrtFree(blockTableDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例代码的调用要点:
- 两段式接口:先调用
aclnnNsaCompressWithCacheGetWorkspaceSize获取workspaceSize与executor;当workspaceSize > 0时用aclrtMalloc在 Device 侧申请 workspace;再调用aclnnNsaCompressWithCache真正下发执行,最后通过aclrtSynchronizeStream同步等待任务完成; - 可选输入按场景传递:PageAttention 场景传
blockTable与actSeqLen(本例两者均传入,actSeqLenHostData[0]=32表示第一个 batch 当前序列长度为 32,恰好满足(32-32)%16==0的压缩触发条件);若使用 TND 排布而非 PageAttention,则需传layout="TND"且blockTable传 nullptr; - 资源管理:示例末尾对 aclTensor、Device 内存、Stream 与 device 依次释放,避免资源泄漏。
仓库测试目录 tests/comm/inc/aclnn_nsa_compress_with_cache_param.h 中的AclnnNsaCompressWithCacheParam以(batchSize, headNum, headDim, dtype, actSeqLenList, layoutType, compressBlockSize, compressStride, actSeqLenType, pageBlockSize)构造参数组合,覆盖了不同 batch、head、压缩滑窗与排布(LayoutType)下的用例,可以作为自建测试时参数设计的参考;对应 kernel/tiling 的算子级单测见 tests/utest 目录。
底层实现原理:tiling 分核与 kernel 广播策略
Tiling 数据结构
tiling 阶段生成的NsaCompressWithCacheTilingData定义于 nsa_compress_with_cache_tiling.h,关键字段包括:
kvCacheSize、weightSize、compressKvCacheSize:全量数据的元素个数,用于建立 Global Memory 张量视图;kvCacheSizePerCore、weightSizePerCore、compressKvCacheSizePerCore:每个核单次处理的元素个数;coresNumPerCompress:一个 token 的压缩需要分多少个核协作完成,用于跨核切片(slice)数据搬运;headsNumPerCore:每个核一次处理的 head 数;coreStartBatchIdxList、coreStartHeadIdxList(长度 50 的数组):记录每个核从哪个 batch 开始遍历 seq_len、从哪个 head 开始处理,kernel 通过GetBlockIdx()获取当前核编号后按此数组切分任务;tokenNumPerTile、tokenRepeatCount:核内循环一次压缩的 token 数,以及压缩完compressBlockSize个 token 需要循环的次数;bufferNum:流水(pipeline)缓冲数量,用于 GM→UB 数据搬运与计算的乒乓重叠;isEmpty:输入是否为空,kernel 入口读取后直接提前返回。
从 nsa_compress_with_cache_tiling.cpp 可以看到 tiling 准备阶段会获取平台 AIV 核数(aivNum)与 UB 内存大小(ubSize),二者为 0 时直接报错返回;这些平台信息决定了上述分核参数的计算。
Kernel 流水线处理流程
kernel 主体定义于 op_kernel/nsa_compress_with_cache.cpp,以基类KernelNsaCompressWithCacheBase<T>为核心,处理流程可归纳为:
- InitBase:按核编号
blockIdx从coreStartHeadIdxList/coreStartBatchIdxList中取本核负责的 head 区间与 batch 区间,并为输入队列、输出队列与 VECCALC 计算缓冲申请 UB 空间(InitBase); - InitActSeqLen:遍历本核负责的 batch,用
actSeqLenGm.GetValue(i)读取每个 batch 的序列长度,筛选满足curSeqLen >= compressBlockSize且(curSeqLen - compressBlockSize) % compressStride == 0的 batch 写入compressBatchIdxList(InitActSeqLen),与文档中算法流程第 1 步完全对应; - InitOffset:根据
slotMappingGm.GetValue(batchIdx)计算输出偏移outputOffset,确定压缩结果在compressKvCacheGm中的写入位置(InitOffset); - CopyKvCache:将待压缩 token 从 GM 搬运到 UB,通过
blockTableGm.GetValue(batchIdx * pageNumPerBatch + curBlockPageIdx)完成逻辑 page 到物理 block 的映射,再按blockLen、srcStride(跨核切片跳读)、blockCount参数做带 stride 的DataCopy(CopyKvCache); - ComputeSubTile + ComputeReduce:KV 数据 Cast 为 FP32 后与广播后的权重逐元素相乘,先按 tile 局部归约,再通过
ReduceBlock将compressBlockSize归约到 1 个 token,最后 Cast 回 FP16/BF16(CAST_RINT舍入)送入输出队列; - CopyOut:将压缩结果从 UB 写回
compressKvCacheGm[outputOffset](CopyOut)。
其中多核协作方式为:coresNumPerCompress个核共同完成一个 token 的压缩,每个核只处理headNum / coresNumPerCompress个 head 的切片(headsNumPerCore),从而把大 headNum 的压缩任务切分到多个核上并行执行。
三种权重广播策略
由于权重weight的 shape 为[compressBlockSize, N](只对 token 维与 head 维有值,headDim 维需广播),kernel 按 headDim 与核心数配置在三种广播策略中选择一种,通过TILING_KEY_IS(0/1/2)分发(见 nsa_compress_with_cache.cpp):
| 策略 | 类 | 适用思路 |
|---|---|---|
| DoubleBroadcast(TILING_KEY=0) | KernelNsaCompressWithCacheDoubleBroadcast | 权重先在 head 切片内做第一次广播,跳读切分 headNum 后,再按 headDim 做第二次广播,最终铺满tokenNumPerTile × headsNumPerCore × headDim形状(InitWeightDoubleBroadCast); |
| TileBroadcast(TILING_KEY=1) | KernelNsaCompressWithCacheTileBroadcast | 跳读切分 headNum 后直接按 headDim 广播(InitWeightTileBroadCast); |
| FullBroadcast(TILING_KEY=2) | KernelNsaCompressWithCacheFullBroadcast | 直接对整块 tile 权重按 headDim 做一次广播(InitWeightFullBroadCast)。 |
从源码结构可以推断,三种策略分别针对不同的 headNum/headDim/核数组合,以平衡 UB 中权重广播缓冲的占用与广播指令的开销,是 kernel 针对不同 shape 场景的性能优化分支。
另外,kernel 入口对平台做了条件编译区分:FP16 实现全平台可用;而 BF16 实现通过#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))排除特定 NPU 架构,与 README 中「Kirin 平台不支持 BFLOAT16」的约束一致。
总结
NsaCompressWithCache是 CANN ops-transformer 中面向 NSA 推理场景的 KV 压缩算子,通过「检查序列长度 → 定位压缩数据 → 乘加压缩 → 按 slot 写回」四步流程,将 PagedAttention 缓存中凑满一个compressBlockSize滑窗的 token 流式压缩为稀疏注意力所需的compress_token。本文完整覆盖了其产品支持矩阵、计算公式、算子级与 aclnn 接口级参数、约束与错误码、可编译调用示例,并结合 op_host 与 op_kernel 源码剖析了 tiling 分核调度、page 到物理 block 的映射以及三种权重广播策略的底层实现。如需进一步验证行为,可参考 tests/comm 下的参数构造与 tests/utest 下的 tiling/kernel 单测用例。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考