CANN ops-transformer 算子解析NsaCompressWithCache 推理阶段 NSA KV 压缩算子实战指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 开源仓库中的NsaCompressWithCache算子展开系统讲解其在 Native-Sparse-AttentionNSA大模型推理场景下压缩 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最后CastCAST_RINT舍入回 FP16/BF16 输出。参数说明算子在底层 GE 图中的输入输出与属性定义见 nsa_compress_with_cache_def.cpp。下表汇总了算子层面对外暴露的全部参数参数名输入/输出/属性描述数据类型数据格式s属性当前 batch 的 token 长度。INT64-compressBlockSize属性压缩滑窗大小。INT64-stride属性两次压缩滑窗间隔大小。INT64-weight输入k/v 值的压缩 weight。BFLOAT16、FLOAT16NDinput输入k/v 值的 cache。BFLOAT16、FLOAT16NDslotMapping输入每个 batch 尾部压缩数据存储的位置的索引。INT32NDoutputCacheRef输入/输出输出的 cache。BFLOAT16、FLOAT16ND需要特别说明的是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非连续 Tensorinput输入待压缩张量。不支持空 Tensor与 weight 满足 broadcast 关系input 的第三维大小与 weight 的第二维大小相等headDim 是 16 的整数倍且 ≤256headNum≤64且 headNum50 时 headNum%20。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≥compressStridecompressBlockSize≤64。INT64---compressStride输入两次压缩间的滑窗间隔大小。仅支持取值 16、32、48、64。INT64---actSeqLenType输入actSeqLenOptional 的不同表达形式。actSeqLenOptional 有输入时生效可取值 0 或 10 代表 actSeqLenOptional 中数值为前继 batch 序列大小的 cumsum累积和结果1 代表其中数值为每个 batch 的序列大小当前仅支持 1。INT64---pageBlockSize输入page attention 场景下 page 的 blocksize 大小。只能是 64 或者 128。INT64---outputCache输出压缩之后的 cache。数据类型与 input 保持一致。FLOAT16、BFLOAT16ND[result_len, N, D]×workspaceSize输出返回需要在 Device 侧申请的 workspace 大小。-----executor输出返回 op 执行器包含了算子计算流程。-----关于输入排布需要区分两种场景PageAttention 场景下传入blockTableOptionalinput的 shape 支持[blockNum, pageBlockSize, N, D]其余场景TND 排布传入actSeqLenOptional与layoutOptionalTNDinput的 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完成算子调度。约束说明使用该算子必须满足以下约束input和weight满足 broadcast 关系input的第三维大小与weight的第二维大小相等compressBlockSize、stride必须是 16 的整数倍且compressBlockSize stridecompressBlockSize 64actSeqLenType目前仅支持取值 1layoutOptional取值可以是 BSH、SBH、BSND、BNSD、TND但当前不会生效pageBlockSize只能是 64 或者 128headDim是 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_size4、headNum24、headDim192、pageBlockSize128、compressBlockSize32、compressStride16、maxSeqLen512演示完整调用#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::vectorint64_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::vectorT hostData, const std::vectorint64_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::vectorint64_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::vectorint64_t inputShape {blocks_num, page_block_size, heads_num, heads_dim}; std::vectorint64_t weightShape {compress_block_size, heads_num}; std::vectorint64_t slotMappingShape {batch_size}; std::vectorint64_t outputCacheRefShape {result_len, heads_num, heads_dim}; std::vectorint64_t actSeqLenShape {batch_size}; std::vectorint64_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::vectoraclFloat16 inputHostData(inputShape[0] * inputShape[1] * inputShape[2] * inputShape[3], aclFloatToFloat16(1.0)); std::vectoraclFloat16 weightHostData(weightShape[0] * weightShape[1], aclFloatToFloat16(1.0)); std::vectorint32_t slotMappingHostData(slotMappingShape[0], 0); std::vectoraclFloat16 outputCacheRefHostData(outputCacheRefShape[0] * outputCacheRefShape[1] * outputCacheRefShape[2], aclFloatToFloat16(1.0)); std::vectorint64_t actSeqLenHostData(actSeqLenShape[0], 0); std::vectorint32_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::vectoraclFloat16 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)%160的压缩触发条件若使用 TND 排布而非 PageAttention则需传layoutTND且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以基类KernelNsaCompressWithCacheBaseT为核心处理流程可归纳为InitBase按核编号blockIdx从coreStartHeadIdxList/coreStartBatchIdxList中取本核负责的 head 区间与 batch 区间并为输入队列、输出队列与 VECCALC 计算缓冲申请 UB 空间InitBaseInitActSeqLen遍历本核负责的 batch用actSeqLenGm.GetValue(i)读取每个 batch 的序列长度筛选满足curSeqLen compressBlockSize且(curSeqLen - compressBlockSize) % compressStride 0的 batch 写入compressBatchIdxListInitActSeqLen与文档中算法流程第 1 步完全对应InitOffset根据slotMappingGm.GetValue(batchIdx)计算输出偏移outputOffset确定压缩结果在compressKvCacheGm中的写入位置InitOffsetCopyKvCache将待压缩 token 从 GM 搬运到 UB通过blockTableGm.GetValue(batchIdx * pageNumPerBatch curBlockPageIdx)完成逻辑 page 到物理 block 的映射再按blockLen、srcStride跨核切片跳读、blockCount参数做带 stride 的DataCopyCopyKvCacheComputeSubTile ComputeReduceKV 数据 Cast 为 FP32 后与广播后的权重逐元素相乘先按 tile 局部归约再通过ReduceBlock将compressBlockSize归约到 1 个 token最后 Cast 回 FP16/BF16CAST_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策略类适用思路DoubleBroadcastTILING_KEY0KernelNsaCompressWithCacheDoubleBroadcast权重先在 head 切片内做第一次广播跳读切分 headNum 后再按 headDim 做第二次广播最终铺满tokenNumPerTile × headsNumPerCore × headDim形状InitWeightDoubleBroadCastTileBroadcastTILING_KEY1KernelNsaCompressWithCacheTileBroadcast跳读切分 headNum 后直接按 headDim 广播InitWeightTileBroadCastFullBroadcastTILING_KEY2KernelNsaCompressWithCacheFullBroadcast直接对整块 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),仅供参考