CANN ops-transformer 算子解析:NsaCompressWithCache 推理阶段 NSA KV 压缩算子实战指南
2026/9/18 14:21:55 网站建设 项目流程

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()kirinx90kirin9030两个配置名注册 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写入压缩缓存。

算法流程(四步)

  1. 检查序列长度:遍历actSeqLenOptional,检查是否存在满足s >= compressBlockSize(s - compressBlockSize) % stride == 0的序列长度(其中s为当前 batch 的 token 长度);
  2. 定位待压缩数据:找到满足序列长度条件的batchIdx,根据blockTableOptional找到该 batch 的后compressBlockSize个 token 作为待压缩数据;
  3. 执行压缩算法:对待压缩 token 与压缩权重执行加权求和(乘加归约);
  4. 写回压缩缓存:根据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,最后CastCAST_RINT舍入)回 FP16/BF16 输出。

参数说明

算子在底层 GE 图中的输入输出与属性定义见 nsa_compress_with_cache_def.cpp。下表汇总了算子层面对外暴露的全部参数:

参数名输入/输出/属性描述数据类型数据格式
s属性当前 batch 的 token 长度。INT64-
compressBlockSize属性压缩滑窗大小。INT64-
stride属性两次压缩滑窗间隔大小。INT64-
weight输入k/v 值的压缩 weight。BFLOAT16、FLOAT16ND
input输入k/v 值的 cache。BFLOAT16、FLOAT16ND
slotMapping输入每个 batch 尾部压缩数据存储的位置的索引。INT32ND
outputCacheRef输入/输出输出的 cache。BFLOAT16、FLOAT16ND

需要特别说明的是:

  • Kirin X90/Kirin 9030 处理器系列产品不支持 BFLOAT16,该平台下input/weight/outputCache仅支持 FLOAT16;
  • 从 nsa_compress_with_cache_def.cpp 可以看到,act_seq_lenblock_table在底层注册为 OPTIONAL(可选)输入,其中act_seq_len还带有ValueDepend(OPTIONAL)标记,表示其数值会参与 tiling 决策;layoutcompress_block_sizecompress_strideact_seq_len_typepage_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、BFLOAT16ND[blockNum, pageBlockSize, N, D]、[TND]
weight输入压缩的权重。不支持空 Tensor;数据类型与 input 保持一致。FLOAT16、BFLOAT16ND[compressBlockSize, N]
slotMapping输入每个 batch 尾部压缩数据存储的位置的索引。不支持空 Tensor;值无重复,否则会导致计算结果不稳定。INT32ND[B]×
actSeqLenOptional可选输入每个 Batch 对应的 S 大小。在 TND 排布场景下需要该输入,其余场景输入 nullptr;值不应超过序列最大长度。S 表示输入样本序列长度。INT64ND[B]-
blockTableOptional可选输入PageAttention 中 KV 存储使用的 block 映射表。不使用该功能可传入 nullptr;值不超过 blockNum,否则会发生越界。INT32ND[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 的序列大小;当前仅支持 1INT64---
pageBlockSize输入page attention 场景下 page 的 blocksize 大小。只能是 64 或者 128。INT64---
outputCache输出压缩之后的 cache。数据类型与 input 保持一致。FLOAT16、BFLOAT16ND[result_len, N, D]×
workspaceSize输出返回需要在 Device 侧申请的 workspace 大小。-----
executor输出返回 op 执行器,包含了算子计算流程。-----

关于输入排布需要区分两种场景:PageAttention 场景下(传入blockTableOptional),input的 shape 支持[blockNum, pageBlockSize, N, D]其余场景(TND 排布,传入actSeqLenOptionallayoutOptional="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_NULLPTR161001计算输入和必选计算输出是空指针。
ACLNN_ERR_PARAM_INVALID161002计算输入和输出的数据类型和格式不在支持的范围内。
ACLNN_ERR_PARAM_INVALID161002input、weight、outputCache 为空 tensor。
ACLNN_ERR_INNER_TILING_ERROR561002input 和 weight 不满足 broadcast 关系,即 input 的第三维大小与 weight 的第二维大小不相等。
ACLNN_ERR_INNER_TILING_ERROR561002activeNum、expertNum、expertCapacity 的值小于 0。
ACLNN_ERR_INNER_TILING_ERROR561002compress_block_size、compress_stride 不是 16 的整数倍。
ACLNN_ERR_INNER_TILING_ERROR561002seq_lens_type != 1,或者 layout 取值不是 BSH、SBH、BSND、BNSD、TND 中的一个。
ACLNN_ERR_INNER_TILING_ERROR561002page_block_size 取值不是 64 或者 128。
ACLNN_ERR_INNER_TILING_ERROR561002headDim 未对齐 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完成算子调度。

约束说明

使用该算子必须满足以下约束:

  • inputweight满足 broadcast 关系,input的第三维大小与weight的第二维大小相等;
  • compressBlockSizestride必须是 16 的整数倍,且compressBlockSize >= stridecompressBlockSize <= 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 > 50headNum % 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=4headNum=24headDim=192pageBlockSize=128compressBlockSize=32compressStride=16maxSeqLen=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获取workspaceSizeexecutor;当workspaceSize > 0时用aclrtMalloc在 Device 侧申请 workspace;再调用aclnnNsaCompressWithCache真正下发执行,最后通过aclrtSynchronizeStream同步等待任务完成;
  • 可选输入按场景传递:PageAttention 场景传blockTableactSeqLen(本例两者均传入,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,关键字段包括:

  • kvCacheSizeweightSizecompressKvCacheSize:全量数据的元素个数,用于建立 Global Memory 张量视图;
  • kvCacheSizePerCoreweightSizePerCorecompressKvCacheSizePerCore:每个核单次处理的元素个数;
  • coresNumPerCompress:一个 token 的压缩需要分多少个核协作完成,用于跨核切片(slice)数据搬运;
  • headsNumPerCore:每个核一次处理的 head 数;
  • coreStartBatchIdxListcoreStartHeadIdxList(长度 50 的数组):记录每个核从哪个 batch 开始遍历 seq_len、从哪个 head 开始处理,kernel 通过GetBlockIdx()获取当前核编号后按此数组切分任务;
  • tokenNumPerTiletokenRepeatCount:核内循环一次压缩的 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>为核心,处理流程可归纳为:

  1. InitBase:按核编号blockIdxcoreStartHeadIdxList/coreStartBatchIdxList中取本核负责的 head 区间与 batch 区间,并为输入队列、输出队列与 VECCALC 计算缓冲申请 UB 空间(InitBase);
  2. InitActSeqLen:遍历本核负责的 batch,用actSeqLenGm.GetValue(i)读取每个 batch 的序列长度,筛选满足curSeqLen >= compressBlockSize(curSeqLen - compressBlockSize) % compressStride == 0的 batch 写入compressBatchIdxList(InitActSeqLen),与文档中算法流程第 1 步完全对应;
  3. InitOffset:根据slotMappingGm.GetValue(batchIdx)计算输出偏移outputOffset,确定压缩结果在compressKvCacheGm中的写入位置(InitOffset);
  4. CopyKvCache:将待压缩 token 从 GM 搬运到 UB,通过blockTableGm.GetValue(batchIdx * pageNumPerBatch + curBlockPageIdx)完成逻辑 page 到物理 block 的映射,再按blockLensrcStride(跨核切片跳读)、blockCount参数做带 stride 的DataCopy(CopyKvCache);
  5. ComputeSubTile + ComputeReduce:KV 数据 Cast 为 FP32 后与广播后的权重逐元素相乘,先按 tile 局部归约,再通过ReduceBlockcompressBlockSize归约到 1 个 token,最后 Cast 回 FP16/BF16(CAST_RINT舍入)送入输出队列;
  6. 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),仅供参考

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

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

立即咨询