拓冰建站拓冰建站
首页 / 资讯中心 / 正文

CANN ops-transformer 算子 aclnnAttentionUpdate 接口详解:SP 域序列并行 Attention 局部结果到全局结果的合并更新

CANN ops-transformer 算子 aclnnAttentionUpdate 接口详解SP 域序列并行 Attention 局部结果到全局结果的合并更新【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer导读aclnnAttentionUpdate是 CANN ops-transformer 算子库attention/attention_update中专门面向序列并行Sequence ParallelSP场景的融合更新算子它将各 SP 域 PagedAttentionPA类算子输出的局部 log-sum-explse与局部 Attention 输出localOut合并为全局结果是序列并行大模型推理/训练中跨域归一化的关键一环。本文以 attention/attention_update/docs/aclnnAttentionUpdate.md 为骨架结合仓库内算子定义、tiling、kernel 与测试源码完整讲解其数学原理、两段式接口原型、参数约束、错误码语义、C 调用示例与底层实现机制帮助开发者快速上手并深入理解该算子在 NPU 上的执行方式。功能定位序列并行下的局部结果合并在长序列超长 context大模型场景中单个 NPU 无法容纳整条序列的 KV Cache通常采用序列并行SP把序列切成多个分片SP 域每个分片由独立的计算单元执行 PagedAttention 等算子。此时每个 SP 域只会得到局部的softmax 统计量lse_ilog-sum-exp和局部的Attention 输出O_i即localOut这些局部结果相互独立不能直接拼接必须先做跨域合并才能得到与整条序列语义等价的全局输出。aclnnAttentionUpdate正是完成这一合并任务的算子它接收sp个 SP 域的局部 lse 与局部 Attention 输出通过求全局最大值 → 重新指数归一化 → 加权求和的流程输出全局 lse可选与全局 Attention 输出即接口功能将各 SP 域 PA 算子的输出的中间结果lse、localOut两个局部变量结果更新成全局结果出自 aclnnAttentionUpdate.md。对应地算子内部输入在 attention_update_def.cpp 中定义为动态个数的lse1 维与go2 维即 localOut张量列表并通过属性sp声明张量个数。数学原理跨 SP 域的 softmax 合并公式设第i个 SP 域的局部统计量为lse_i、局部输出为O_ii 1 … sp算子按如下 4 步将局部量合并为全局量$$ lse_{max} \text{max}_i, lse_i $$$$ lse \sum_i \text{exp}(lse_i - lse_{max}) $$$$ lse_m lse_{max} \text{log}(lse) $$$$ O \sum_i O_i \cdot \text{exp}(lse_i - lse_m) $$其中lse_max是所有 SP 域局部 lse 的最大值用于数值稳定性防止指数溢出lse是重归一化后的指数和lse_m是合并后的全局log-sum-exp即最终的全局归一化常数O是加权求和得到的全局 Attention 输出每个局部输出O_i的权重为exp(lse_i - lse_m)。这一过程在数学上等价于对多个子序列 softmax 结果做log-sum-exp 形式的精确合并不损失精度。仓库测试中的 CPU 参考实现 executor_aclnnAttentionUpdate.py 也严格按照该公式构造先torch.exp再按总和归一化、加权求和用于对 NPU 结果做一致性比对可作为公式理解的辅助参考。产品支持情况根据文档 aclnnAttentionUpdate.md各产品形态支持情况如下产品是否支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持与算子定义源码中AICore().AddConfig注册的硬件配置一致ascend910bAtlas A2 系列、ascend910_93Atlas A3 系列、ascend950三个平台见 attention_update_def.cpp其中ascend950走独立的 regbasearch35实现路径。两段式接口先获取 workspace再执行计算aclnnAttentionUpdate采用 CANN 单算子 API 通用的两段式接口机制说明见 docs/zh/context/two_phase_api.md先调用aclnnAttentionUpdateGetWorkspaceSize完成入参校验并获取计算所需 workspace 大小及封装了计算流程的执行器按返回的workspaceSize在 Device 侧申请内存后再调用aclnnAttentionUpdate执行计算。两段接口的函数原型如下出自 aclnnAttentionUpdate.mdaclnnStatus aclnnAttentionUpdateGetWorkspaceSize( const aclTensorList *lse, const aclTensorList *localOut, int64_t updateType, aclTensor *out, aclTensor *lseOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnAttentionUpdate( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)注意第二段接口aclnnAttentionUpdate不可重复调用同一executor只能执行一次计算否则行为异常。aclnnAttentionUpdateGetWorkspaceSize 参数详解第一段接口共 7 个入参/出参完整参数说明如下整理自 aclnnAttentionUpdate.md参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续 Tensorlse输入各 SP 域的局部 lsetensorList 长度为 spFLOAT32ND[batch * seqLen * headNum]xlocalOut输入各 SP 域的局部 attention outtensorList 长度为 spFLOAT32FLOAT16BFLOAT16ND[batch * seqLen * headNum, headDim]xupdateType输入控制 lseOut 是否输出支持 0、1分别表示不输出 lseOut、输出 lseOutINT64---out输出输出的 tensor-与 localOut 一致ND[batch * seqLen * headNum, headDim]xlseOut可选输出作为 lse_m 可选输出不输出 lseOut 可传入 nullptrFLOAT32ND[batch * seqLen * headNum]xworkspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----executor输出返回 op 执行器包含了算子计算流程-----补充说明lse / localOut 的 tensorList 长度必须等于 sp且两个列表长度必须一致lse[i]与localOut[i]一一对应同一个 SP 域localOut支持 FLOAT32/FLOAT16/BFLOAT16而lse恒为 FLOAT32softmax 统计量统一用高精度 FP32 表达lseOut同样为 FLOAT32第一维batch * seqLen * headNum即文档与源码中的bsh语义batch × seq × headNum它也是 tiling 分核的基本单位非连续 Tensor列标记为 x表示不支持非连续张量源码中第一段接口会通过l0op::Contiguous对每个输入张量做连续性规整后再送入计算图见 aclnn_attention_update.cppupdateType与lseOut存在强绑定关系详见下文错误码部分。返回值与错误码语义两段接口均返回aclnnStatus通用返回码说明见 docs/zh/context/aclnn_return_code.md。第一段接口aclnnAttentionUpdateGetWorkspaceSize完成入参校验出现以下场景时报错整理自 aclnnAttentionUpdate.md返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入的 lse、localOut 或者 out 是空指针ACLNN_ERR_PARAM_INVALID161002传入的 lse、localOut 或者 out 的数据类型/数据格式不在支持的范围之内ACLNN_ERR_PARAM_INVALID161002传入的 updateType 或者 sp 不在取值范围ACLNN_ERR_PARAM_INVALID161002传入的 lse、localOut 或者 out 的 shape 不满足约束ACLNN_ERR_PARAM_INVALID161002updateType 为 0 时传入的 lseOut 不为 nullptrACLNN_ERR_PARAM_INVALID161002updateType 为 1 时传入的 lseOut 为 nullptr这些校验逻辑在源码中有完整的对应实现空指针检查CheckNotNull逐元素校验lse、localOut列表内的每个 tensor 以及outaclnn_attention_update.cpp数据类型检查CheckDtypeValid校验 lse 为 FLOAT、localOut/out 在 FLOAT/FLOAT16/BF16 范围内aclnn_attention_update.cppAscend 950 的 regbase 分支CheckDtypeValid_95还额外要求 localOut 列表内各张量 dtype 一致、out 与 localOut[0] 一致aclnn_attention_update.cppsp 范围与列表长度一致性CheckSpsp ∈ [1, 128]与CheckSp_95Ascend 950 上 sp ∈ [1, 16]aclnn_attention_update.cppshape 约束CheckShape/CheckShape_95校验 lse 为 1 维、localOut/out 为 2 维、所有张量 shape 相同、localOut[1]为 8 的倍数且 ≤ 512、各张量第一维bsh一致、updateType 1时 lseOut 为 1 维且第一维一致aclnn_attention_update.cppupdateType 与 lseOut 绑定关系CheckUpdateTypeAndLseOutaclnn_attention_update.cpp。约束说明使用aclnnAttentionUpdate需遵守以下约束整理自 aclnnAttentionUpdate.md确定性计算aclnnAttentionUpdate默认为确定性实现相关概念见 docs/zh/context/determinism_compute.md相同输入可复现相同结果sp 取值范围Atlas A2 训练系列产品 / Atlas A2 推理系列产品、Atlas A3 训练系列产品 / Atlas A3 推理系列产品[1, 128]Ascend 950PR / Ascend 950DT[1, 16]headDim 取值范围[8, 512]且是 8 的倍数支持空 Tensor当lse[0]或localOut[0]为空张量时第一段接口直接返回workspaceSize 0并跳过计算见 aclnn_attention_update.cpp不支持非连续 Tensor如上文所述接口内部会先做Contiguous处理。上述约束与 tiling 源码中的硬校验一一对应attention_update_tiling.cpp中定义了D_MIN 8、D_MAX 512、D_DIVIDE_8 8、ATTR_SP_MAX 16并在CheckInputParamsupdateType ∈ {0,1}、sp ∈ [1,16]、CheckInputDimH 维在 [8,512] 且为 8 的倍数、所有 lse/go 首维一致、CheckInputDtypego ∈ {FLOAT, FLOAT16, BF16}lse 恒为 FLOAT中执行见 attention_update_tiling.cpp。tiling 层 sp 上限取 16而 aclnn 接口层在 A2/A3 上放宽到 128两处范围差异由接口层的连续化与计算图拼接逻辑多个l0op::AttentionUpdate组合衔接。完整调用示例与逐步讲解文档给出了完整的可直接参考的调用示例见 aclnnAttentionUpdate.md仓库中另有可独立编译运行的简化版样例 examples/test_aclnn_attention_update.cpp。整体调用流程含编译运行环境的准备请参考 docs/zh/context/compile_and_run_sample.md如下1. 环境初始化固定写法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; }依次完成aclInit、aclrtSetDevice、aclrtCreateStream这是所有 aclnn 单算子调用的固定前置步骤。2. 构造输入与输出以sp 2、bsh 256、headDim 128为例std::vectorint64_t lseShape {256}; // [batch * seqLen * headNum] std::vectorint64_t localOutShape {256, 128}; // [batch * seqLen * headNum, headDim] std::vectorint64_t outShape {256, 128}; int64_t updateType 0; // 0不输出 lseOut1输出 lseOut void* lseDeviceAddr[2] {nullptr, nullptr}; void* localOutDeviceAddr[2] {nullptr, nullptr}; void* outDeviceAddr nullptr; std::vectoraclTensor* lse {nullptr, nullptr}; std::vectoraclTensor* localOut {nullptr, nullptr}; aclTensor* out nullptr;对每个张量调用aclrtMalloc申请 Device 侧内存、aclrtMemcpy拷入 host 数据再通过aclCreateTensor创建aclTensor连续张量需按 shape 计算 stridetemplate 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); auto ret aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, ...); ret aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, ...); std::vectorint64_t stride(shape.size(), 1); for (int64_t i shape.size() - 2; i 0; i--) { stride[i] shape[i 1] * stride[i 1]; } *tensor aclCreateTensor(shape.data(), shape.size(), dataType, stride.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; }由于lse与localOut是 tensorList需分别调用aclCreateTensorList将张量指针数组封装为aclTensorListaclTensorList *lseList aclCreateTensorList(lse.data(), lse.size()); aclTensorList *localOutList aclCreateTensorList(localOut.data(), localOut.size());3. 两段式调用先调用第一段接口获取 workspaceSize 与 executor注意updateType 0时lseOut传nullptruint64_t workspaceSize 0; aclOpExecutor* executor; ret aclnnAttentionUpdateGetWorkspaceSize(lseList, localOutList, updateType, out, nullptr, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnAttentionUpdateGetWorkspaceSize failed. ERROR: %d\n, ret); return ret);按返回大小申请 workspace 后调用第二段接口void* workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, ...); } ret aclnnAttentionUpdate(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnAttentionUpdate failed. ERROR: %d\n, ret); return ret);4. 同步与结果回拷ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, ...); auto size GetShapeSize(outShape); std::vectorfloat outData(size, 0); ret aclrtMemcpy(outData.data(), outData.size() * sizeof(outData[0]), outDeviceAddr, size * sizeof(outData[0]), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, ...); for (int64_t i 0; i size; i) { LOG_PRINT(out result[%ld] is: %f\n, i, outData[i]); }最后通过aclrtDestroyStream、aclrtResetDevice、aclFinalize完成资源释放完整释放逻辑见示例代码Finalize函数。源码级实现原理算子图定义在 attention_update_def.cpp 中输入lseDYNAMICFLOAT/ND与goDYNAMICFLOAT/FLOAT16/BF16/ND为动态个数的 tensorList个数由属性sp决定输出output与 go 同 dtype与lse_mFLOAT属性update_type默认 0与sp必填ascend950配置走attention_update_aptregbase/动态编译路径见ExtendCfgInfo(opFile.value, attention_update_apt)并开启DynamicRankSupportFlag、DynamicShapeSupportFlag、PrecisionReduceFlag。接口层组装第一段接口aclnnAttentionUpdateGetWorkspaceSize在完成校验后将sp个输入张量逐个Contiguous规整组装为 tensorList交给l0op::AttentionUpdate构造计算图并用l0op::ViewCopy将中间结果拷贝到用户提供的out/lseOut张量最终通过uniqueExecutor-GetWorkspaceSize()返回所需 workspace 大小见 aclnn_attention_update.cpp。Tiling 策略attention_update_tiling.cpp实现了基于 bshbatch × seq × headNum维度的多核切分attention_update_tiling.cpp按CeilDiv(bshSize, totalCoreNum)计算每核处理量perCoreCount尾核处理lastCoreCount依据 UB 容量与 double bufferDOUBLE_BUFFER_NUM 2开销计算单次内循环可承载的 bsh 数bshInLoop其中输入因子为sp × lse sp × dAlign × go的双缓冲占用中间计算用 FP32 提升精度输出同样开双缓冲预留sp * ubBlockSize空间防止 lse 对齐搬入 UB 后占用膨胀通过context-SetBlockDim(usedCoreNum)设置核数workspace 固定申请16 * 1024 * 1024字节SYS_WORKSPACE_SIZE见 attention_update_tiling.cpp。tiling data 中完整记录 sp、d、usedCoreNum、perCoreCount、lastCoreCount、perCoreLoops、lastCoreLoops、perCorePerLoopCount、bshInLoop 等切分参数定义见 attention_update_tiling.h并按 tiling key20000 updateType注册空张量场景为 10000。Kernel 执行NPU kernel 入口在 op_kernel/attention_update.cpp仅使用 AIVKERNEL_TYPE_AIV_ONLY按 tiling key 实例化不同精度的DecodeUpdatefloat, goType模板执行。Ascend 950 的 regbase 实现arch35/attention_update_with_lse_regbase.h展示了核心计算过程通过ListTensorDesc读取 sp 个 lse/go 全局张量地址ComputeMaxVF用向量寄存器对 sp 个 lse 逐元素求Max将INF先替换为-INF做保护随后逐域做Sub → Exp → Add累加、Log得到全局lse_m并回写ComputeOutVF按UNROLL_NUM 2展开 sp 域用MulAddDst完成O_i × exp(lse_i - lse_m)的加权累加得到全局输出 O数据搬入搬出全部采用双缓冲BUFFER_NUM 2与DataCopyPad对齐d 维按 UB block 对齐CeilAlign(d, goBlockNum)。测试验证仓库为该算子提供了完善的 ST/UT 覆盖可用于验证正确性与精度ST 用例atk_aclnnAttentionUpdate.json中配置了sp 2、bsh 20480、headDim 512、updateType 0的边界用例要求high_precision精度见 tests/st/aclnnAttentionUpdate/atk_aclnnAttentionUpdate.jsonCPU 参考实现executor_aclnnAttentionUpdate.py用 PyTorch 严格按合并公式实现 reference供 NPU 结果比对executor_aclnnAttentionUpdate.pybf16/fp16 专项tests/st/aclnnAttentionUpdate_bf16fp16/目录覆盖半精度 localOut 场景UT 测试tests/ut/下分别对 tilingtest_attention_update_tiling.cpp、kerneltest_attention_update.cpp、aclnn 接口层test_aclnn_attention_update.cpp提供单元测试。小结aclnnAttentionUpdate是 CANN ops-transformer 中面向序列并行 Attention 的收尾算子通过 log-sum-exp 形式的跨 SP 域合并将各域的局部 lse 与局部输出精确融合为全局结果。使用时需牢记lse 恒为 FP32、localOut 支持 FP16/BF16/FP32headDim 必须为 [8, 512] 内 8 的倍数sp 上限在 A2/A3 上为 128、在 Ascend 950 上为 16updateType与lseOut的传参必须严格配对。掌握两段式调用流程与上述约束后即可在序列并行推理/训练链路中正确接入该算子实现跨域 Attention 结果的全局归一化合并。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门