尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

CANN ops-transformer 之 NsaCompress 算子实战:基于 NSA 算法的 long-context KV 序列压缩与 aclnnNsaCompress 调用指南

CANN ops-transformer 之 NsaCompress 算子实战:基于 NSA 算法的 long-context KV 序列压缩与 aclnnNsaCompress 调用指南 CANN ops-transformer 之 NsaCompress 算子实战基于 NSA 算法的 long-context KV 序列压缩与 aclnnNsaCompress 调用指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerNsaCompress 是 CANN ops-transformer 注意力算子库中面向训练场景的 KV 压缩算子其核心思想源自 Native Sparse AttentionNSA算法在注意力计算之前先对 K 序列按滑窗做压缩从而显著减轻 long-context 场景下注意力计算的规模。本文以 attention/nsa_compress/README.md 及其配套的 aclnnNsaCompress 接口文档 为主体结合仓库内的算子定义、Shape 推导、Tiling 切分与 Kernel 实现源码系统讲解 NsaCompress 的功能原理、参数与约束、两段式 aclnn 调用流程以及可运行的完整示例帮助你在 Atlas A2 / A3 训练与推理系列产品上快速落地该算子。产品支持情况NsaCompress 算子对昇腾产品的支持情况如下与 README 保持一致产品是否支持Ascend 950PR / Ascend 950DT×Atlas A3 训练系列产品 / Atlas A3 推理系列产品√Atlas A2 训练系列产品 / Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×从算子注册代码 nsa_compress_def.cpp 可以印证该算子通过AICore().AddConfig(ascend910b)与AICore().AddConfig(ascend910_93)分别对应 Atlas A2910B 系列与 Atlas A3 两个算力平台其余平台未注册配置因此不被支持。功能说明与计算公式NsaCompress 算子的功能是在训练场景下使用 NSA Compress 算法减轻 long-context 的注意力计算实现在 KV 序列维度进行压缩。它作用于 K或 KV序列通过一个可学习的压缩权重对滑窗内的连续 token 进行加权聚合得到一个长度显著缩短的压缩序列后续注意力计算将基于该压缩序列进行从而降低 long-context 场景的计算与访存开销。NsaCompress 正向计算公式如下$$ \tilde{K}t^{\text{cmp}} f_K^{\text{cmp}}(k{:t}) \left{ \varphi(k_{id1:idl}) \bigg| 0 \leq i \leq \left\lfloor \frac{t-l}{d} \right\rfloor \right} $$其中l即参数表中的compressBlockSize压缩滑窗大小d即compressStride两次压缩滑窗的间隔φ为基于weight的压缩映射滑动窗口内 token 与压缩权重的加权聚合。该公式说明压缩结果是一系列滑窗输出 token 的集合每个输出 token 对应一个长度为compressBlockSize的输入窗口相邻窗口起点相隔compressStride。输出 token 数量推导由上述公式可以直接推导出单个 batch 的输出 token 数量。在 Shape 推导实现 nsa_compress_infershape.cpp 中for (size_t i 0; i batchSize; i) { int64_t cur_seq_len actualSeqLen[i] - preSeqLen; if (cur_seq_len compressBlockSize) { compressKvNum (cur_seq_len - compressBlockSize compressStride) / compressStride; } preSeqLen cur_seq_len; } gert::Shape out gert::Shape({compressKvNum, headNums, headDims});即只有当当前 batch 的序列长度cur_seq_len compressBlockSize时才产生压缩输出单个 batch 的输出 token 数为(cur_seq_len - compressBlockSize compressStride) / compressStride向上取整等价于(cur_seq_len - compressBlockSize) / compressStride 1所有 batch 累加得到输出的第一维T而输出 shape 仍保持[T, N, D]。参数说明NsaCompress 算子的输入、输出与属性参数如下表所示参数名输入/输出/属性描述数据类型数据格式input输入待压缩张量shape 支持 [T, N, D]。BFLOAT16、FLOAT16NDweight输入压缩权重shape 支持 [compressBlockSize, N]与 input 满足 broadcast 关系。BFLOAT16、FLOAT16NDactSeqLenOptional输入每个 Batch 对应的 S 大小。INT64NDlayoutOptional输入输入数据排布格式支持 BSH、SBH、BSND、BNSD、TND当前仅支持 TND。String-compressBlockSize输入压缩滑窗大小。INT64-compressStride输入两次压缩滑窗间隔大小。INT64-actSeqLenType输入序列长度类型0 表示 cumsum 结果1 表示每个 batch 序列大小当前仅支持 0。INT64-output输出压缩后的结果shape 支持 [T, N, D]。BFLOAT16、FLOAT16ND维度语义说明input数据排布格式支持从多种维度解读其中BBatch输入样本批量大小SSeq-Length输入样本序列长度HHead-Size隐藏层大小NHead-Num多头数DHead-Dim隐藏层最小单元尺寸满足D H / NTB 和 S 合轴紧密排列的数据每个 batch 的 actSeqLen 前缀和累加得到满足T ΣS。约束说明使用 NsaCompress 算子前需要满足以下约束其中大部分校验逻辑可直接在 nsa_compress_tiling.cpp 的CheckParams中找到对应实现该接口与 PyTorch 配合使用时需要保证 CANN 相关包与 PyTorch 相关包的版本匹配。input 和 weight 需要满足 broadcast 关系input.shape[1] weight.shape[1]不支持 input、weight 为空输入。actSeqLenType目前仅支持取值 0即actSeqLenOptional需要是前缀和cumsum模式Tiling 侧会逐项检查前缀和数值单调不减且要求最后一个值等于input.shape[0]。actSeqLenOptional目前不支持为空且必须是一维张量。layoutOptional目前仅支持 TND此时input.shape[0]必须等于actSeqLenOptional[-1]。input.shape[1] weight.shape[1]即 N需要小于等于 128。input.shape[2]即 D必须是 16 的倍数上限 256。weight.shape[0] compressBlockSize必须是 16 的倍数上限 128。compressStride必须是 16 的整数倍并且compressBlockSize compressStride。从源码看Tiling 阶段的CheckParams会对上述约束逐一校验如input.shape[2] % 16 ! 0、compressBlockSize % 16 ! 0、compressStride % 16 ! 0、compressBlockSize compressStride、actseqlenType ! 0等任一不满足都会以OPS_REPORT_VECTOR_INNER_ERR上报错误并终止执行IsEmptyInput还会拒绝input/weightshape size 为 0 的空输入。调用方式与两段式接口NsaCompress 支持通过 aclnn 接口方式调用。示例工程位于 examples/test_aclnn_nsa_compress.cpp调用方式汇总如下调用方式调用样例说明aclnn 调用test_aclnn_nsa_compress.cppTND 场景通过 aclnnNsaCompress 接口方式调用 NsaCompress 算子。与 CANN 算子库中的其他算子一致aclnnNsaCompress 采用两段式接口详见 两段式接口说明必须先调用aclnnNsaCompressGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器executor再调用aclnnNsaCompress执行计算。第一段接口aclnnNsaCompressGetWorkspaceSizeaclnnStatus aclnnNsaCompressGetWorkspaceSize( const aclTensor *input, const aclTensor *weight, const aclIntArray *actSeqLenOptional, char *layoutOptional, int64_t compressBlockSize, int64_t compressStride, int64_t actSeqLenType, aclTensor *output, uint64_t *workspaceSize, aclOpExecutor **executor)该接口的参数说明如下参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续 Tensorinput输入表示待压缩张量。不支持空 Tensor数据类型与 weight 一致shape 支持 [T, N, D]。FLOAT16、BFLOAT16ND3√weight输入表示压缩权重。不支持空 Tensor数据类型与 input 一致与 input 的 shape 满足 broadcast 关系。FLOAT16、BFLOAT16ND2√actSeqLenOptional输入描述每个 Batch 对应的 S 大小。当前不能为空。INT64ND1×layoutOptional输入代表输入 input 的数据排布格式。支持 BSH、SBH、BSND、BNSD、TND当前仅支持 TND。String---compressBlockSize输入压缩滑窗大小。-INT64---compressStride输入两次压缩滑窗间隔大小。-INT64---actSeqLenType输入描述 actSeqLenOptional 数值类型。可取值 0 或 10 表示数值为 cumsum 结果1 表示数值为每个 batch 序列大小当前仅支持 0。INT64---output输出压缩后的结果。不支持空 Tensor数据类型与 input 保持一致shape 支持 [T, N, D]。FLOAT16、BFLOAT16ND3√workspaceSize输出返回需要在 Device 侧申请的 workspace 大小。-----executor输出返回 op 执行器包含算子计算流程。-----第一段接口会完成入参校验出现以下场景时返回对应错误码aclnn 返回码的完整说明参见 aclnn 返回码返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入 input、weight、actSeqLenOptional 或 output 是空指针。ACLNN_ERR_PARAM_INVALID161002input 和 weight 的数据类型不在支持的范围之内。ACLNN_ERR_PARAM_INVALID161002input 和 weight 的 shape 无法做 broadcast。ACLNN_ERR_PARAM_INVALID161002layoutOptional 不合法。第二段接口aclnnNsaCompressaclnnStatus aclnnNsaCompress( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址。workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口aclnnNsaCompressGetWorkspaceSize获取。executor输入op 执行器包含了算子计算流程。stream输入指定执行任务的 Stream。关于确定性计算aclnnNsaCompress 默认为确定性实现确定性计算的背景可参考 确定性计算说明。完整调用示例以下示例代码取自仓库 examples/test_aclnn_nsa_compress.cpp演示了在 TND 布局、前缀和模式actSeqLenType0下调用 NsaCompress 的完整流程。示例中配置为compressBlockSize 32、compressStride 32、batchSize 1、sampleLen 64、headNum 4、headDim 32即输入 shape 为 [64, 4, 32]权重 shape 为 [32, 4]。具体编译和执行过程请参考 编译与运行样例。#include iostream #include vector #include acl/acl.h #include aclnnop/aclnn_nsa_compress.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::vectorint64_t shape) { int64_t shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } void PrintOutResult(std::vectorint64_t shape, void **deviceAddr) { auto size GetShapeSize(shape); std::vectoraclFloat16 resultData(size, 0); auto ret aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), *deviceAddr, 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); for (int64_t i 0; i size; i) { LOG_PRINT(mean result[%ld] is: %f\n, i, aclFloat16ToFloat(resultData[i])); } } int Init(int32_t deviceId, aclrtContext *context, 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 aclrtCreateContext(context, deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateContext failed. ERROR: %d\n, ret); return ret); ret aclrtSetCurrentContext(*context); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetCurrentContext 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); // 调用aclCreateTensor接口创建aclTensor *tensor aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return ACL_SUCCESS; } int main() { // 1.固定写法device/context/stream初始化参考AscendCL对外接口列表 // 根据自己的实际device填写deviceId int32_t deviceId 0; aclrtContext context; aclrtStream stream; auto ret Init(deviceId, context, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 构造输入与输出需要根据API的接口自定义构造 void *inputDeviceAddr nullptr; void *weightDeviceAddr nullptr; void *outputDeviceAddr nullptr; aclTensor *input nullptr; aclTensor *weight nullptr; aclIntArray *actSeqLenOptional nullptr; aclTensor *output nullptr; // 自定义输入与属性 int64_t compressBlockSize 32; int64_t compressStride 32; int64_t actSeqLenType 0; // 0是前缀和模式1是count计数模式 char *layout TND; int32_t batchSize 1; int32_t sampleLen 64; int32_t headNum 4; int32_t headDim 32; std::vectorint64_t inputShape {batchSize * sampleLen, headNum, headDim}; std::vectorint64_t weightShape {compressBlockSize, headNum}; std::vectorint64_t actSeqShape {batchSize}; std::vectoraclFloat16 inputHostData(batchSize * sampleLen * headNum * headDim); std::vectoraclFloat16 weightHostData(compressBlockSize * headNum); std::vectorint64_t actSeqHostData(batchSize); for (int i 0; i inputHostData.size(); i) { inputHostData[i] aclFloatToFloat16(1.0); } for (int i 0; i weightHostData.size(); i) { weightHostData[i] aclFloatToFloat16(1.0); } int outputNum 0; int preActSeqLen 0; for (int i 0; i batchSize; i) { if (actSeqLenType 0) { actSeqHostData[i] sampleLen preActSeqLen; preActSeqLen actSeqHostData[i]; } else if (actSeqLenType 1) { actSeqHostData[i] sampleLen; } if (sampleLen compressBlockSize) { outputNum (sampleLen - compressBlockSize) / compressStride 1; } } std::vectorint64_t outputShape {outputNum, headNum, headDim}; std::vectoraclFloat16 outputHostData(outputNum * headNum * headDim); 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); actSeqLenOptional aclCreateIntArray(actSeqHostData.data(), actSeqHostData.size()); ret CreateAclTensor(outputHostData, outputShape, outputDeviceAddr, aclDataType::ACL_FLOAT16, output); CHECK_RET(ret ACL_SUCCESS, return ret); // 3. 调用CANN算子库API需要修改为具体的Api名称 uint64_t workspaceSize 0; aclOpExecutor *executor; // 调用aclnnNsaCompressGetWorkspaceSize第一段接口 ret aclnnNsaCompressGetWorkspaceSize(input, weight, actSeqLenOptional, layout, compressBlockSize, compressStride, actSeqLenType, output, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnNsaCompressGetWorkspaceSize 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); } // 调用aclnnNsaCompress第二段接口 ret aclnnNsaCompress(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnNsaCompress 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的接口定义修改 PrintOutResult(outputShape, outputDeviceAddr); // 6. 释放aclTensor和aclScalar需要根据具体API的接口定义修改 aclDestroyTensor(input); aclDestroyTensor(weight); aclDestroyIntArray(actSeqLenOptional); aclDestroyTensor(output); // 7. 释放device资源 aclrtFree(inputDeviceAddr); aclrtFree(weightDeviceAddr); // aclrtFree(actSeqDeviceAddr); aclrtFree(outputDeviceAddr); if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtDestroyContext(context); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例中值得注意的几点输出 shape 自算outputNum (sampleLen - compressBlockSize) / compressStride 1当sampleLen compressBlockSize时与算子 Shape 推导逻辑一致本示例中sampleLen64、compressBlockSize32、compressStride32因此输出outputNum 2即输出 shape 为 [2, 4, 32]。前缀和构造actSeqLenType 0时actSeqHostData填入前缀和[64]若为多 batch则依次累加各 batch 的序列长度。数据初始化示例将 input 与 weight 全部初始化为 1.0float16便于直接肉眼校验压缩结果每个输出 token 为对应滑窗内 32 个 1.0 与权重 1.0 的加权聚合结果为滑窗长度。workspace 按需申请仅当第一段接口返回的workspaceSize 0时才调用aclrtMalloc申请并在结束时对应释放。源码级原理从 OpDef 到 Kernel算子定义OpDefnsa_compress_def.cpp 中通过OpDef注册了算子的输入输出与属性input/weight/output均为REQUIRED且数据类型限定为DT_FLOAT16、DT_BF16、格式为FORMAT_NDactSeqLenOptional为OPTIONAL且标记了ValueDepend(OPTIONAL)——这意味着其数值会参与编译期/Tiling 期决策这也解释了为何 Tiling 阶段必须读取该 Tensor 的实际数据来计算输出长度。四个属性layoutOptional默认 TND、compressBlockSize、compressStride、actSeqLenType均为Int/String类型。算力侧的 Shape 推导与 TilingShape 推导nsa_compress_infershape.cpp 根据input.shape取 N、D、actSeqLenOptional取 batchSize 与各前缀和以及属性compressBlockSize、compressStride动态推导输出 shape[compressKvNum, headNums, headDims]并校验compressStride ! 0输出数据类型直接继承 input 的数据类型。Tiling 参数校验nsa_compress_tiling.cpp 的CheckParams实现前述全部约束校验并通过TilingInputsDataDependency({ACT_SEQ_LEN_INPUT_INDEX})声明对actSeqLenOptional数据的依赖。Tiling 切分nsa_compress_tiling_general.cpp 中的NsaCompressTiling完成核心的并行切分策略依据aivNum可用 AI Vector 核数的因子将 N 个 head 划分为divHeadNum组、将输出 token 划分为divSeqNum段为每个核分配 head 索引、输出 token 数量与起始偏移依据输入数据类型fp16/bf16与计算数据类型内部统一转为 fp32 计算估算 UB 内存占用weight 的 fp16/fp32/广播副本、overlap 中间结果、压缩输出等在满足 UB 容量约束下求解每核每次搬运的 KV token 数maxCopyKVTokensNums计算每个核所需搬运的首个 KV 块在输入中的 seq 索引、所属 batch 及全局偏移最终写入NsaCompressTilingData并设置SetBlockDim(aivNum)完成多核并行调度。Kernel 执行流Kernel 入口 nsa_compress.cpp 中通过KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY)强制指定算子类型为 AIV_ONLY并根据输入数据类型ORIG_DTYPE_INPUT实例化KernelNASCompressbfloat16_t或KernelNASCompresshalf。核心类 nsa_compress_kernel.h 中的Process()采用经典的 CopyIn → Compute → CopyOut 流水当当前 sample 剩余可拷贝序列长度不为 0 时持续消费否则切换到下一个 batch 并重置 overlap 上下文直到该核负责的输出 token 全部产出。setTiling从NsaCompressTilingData中恢复全局信息batchSize、headNum、headDim、compressBlockSize、compressStride、maxOverlapNum 等与核内信息起始 batch/seq 索引、负责的 head 数量与索引、输出数量与偏移等。测试验证仓库在 tests/ut/op_host/test_nsa_compress_infershape.cpp 中提供了基于 gtest 的 Shape 推导单测。例如用例NsaCompress_infershape_A1_fp16输入 shape 为 [48, 32, 16]fp16actSeqLen {16, 32, 48}三个 batch 各 16 的前缀和compressBlockSize 16、compressStride 16期望输出 shape 为 [3, 32, 16]——每个 batch 各产出 1 个压缩 token验证了多 batch 前缀和模式下输出长度的累加逻辑NsaCompress_infershape_A2_bf16则以 bf16 验证相同场景。此外还有针对各平台arch22即 Atlas A2 对应架构的 Tiling 单测tests/ut/op_host/arch22/test_nsa_compress_tiling.cpp与 op_api 层调用测试tests/ut/op_api/test_aclnn_nsa_compress.cpp可供你了解算子的预期行为并在本地复跑验证。小结NsaCompress 是 CANN ops-transformer 中面向 long-context 训练场景的 KV 压缩算子它以compressBlockSize与compressStride为滑窗参数、以weight为压缩权重将 K 序列在序列维度压缩为更短的表示。使用前需重点核对产品平台Atlas A2/A3、TND 布局、前缀和模式的actSeqLenOptional、以及 16 字节对齐相关的 shape/参数约束调用侧遵循aclnnNsaCompressGetWorkspaceSizeaclnnNsaCompress的两段式流程仓库提供的 调用示例 可以直接作为开发起点。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表