PyPTO 在线 Softmax 状态更新算子 `online_softmax_update` 使用详解 📅 发布时间:2026/9/19 0:08:48 👁 浏览次数: PyPTO 在线 Softmax 状态更新算子online_softmax_update使用详解【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读pypto.experimental.online_softmax_update是 CANN PyPTOParallel Tensor/Tile Operation 编程范式提供的在线 Softmax 状态更新算子用于在 FlashAttention 等分块注意力场景中把当前 scores 块的局部最大值、指数和与未归一化输出合并进历史累计状态。本文以 官方 API 文档 为主体结合仓库内 Python 前端实现、算子实现 与 TileOp 内核 等源码完整讲解其函数原型、参数语义、约束条件、TileShape 设置以及 FlashAttention 实战用法。读完本文你将掌握该定制接口的调用方式、在线 Softmax 合并公式的底层计算序列以及它与pypto.experimental.online_softmax的配套协作模式。产品支持情况该接口为定制接口仅在下述昇腾产品上受支持其余产品调用会失败Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持这一限制也在代码中有直接体现仓库的算子实现通过CheckSupportedNPUArch(ONLINE_SOFTMAX_SUPPORTED_ARCHITECTURES, OnlineSoftmaxUpdate)做架构校验见 operation_impl.cppST 冒烟测试也统一打上了pytest.mark.soc(950)标记仅面向 950 系列 SoC 运行见 test_online_softmax.py。功能说明在线 Softmax 的状态合并在线onlineSoftmax 是长序列分块注意力中避免为整行 scores 保留全量中间结果的标准技巧先逐块计算局部统计量再增量合并。该接口就是这一流程中的合并环节它完成的工作是给定历史块previous的列最大值、列指数和、未归一化中间输出以及当前块current的列最大值、列指数和、未归一化中间输出按在线 Softmax 公式合并两部分状态输出更新后的最大值、指数和与未归一化输出。假设new_max max(previous_max, current_max)则合并语义可写为updated_max max(previous_max, current_max)updated_sum previous_sum * exp(previous_max - updated_max) current_sum * exp(current_max - updated_max)updated_output previous_output * exp(previous_max - updated_max) current_output * exp(current_max - updated_max)这一公式在 softmax.h 的TOnlineSoftmaxUpdate内核中有逐指令的完整对应先用pto::TMAX计算updatedMax再用TSUBTEXP分别求出历史块与当前块的缩放因子exp(previousMax - updatedMax)、exp(currentMax - updatedMax)随后TMULTADD合并指数和TCOLEXPANDMULTADD完成对未归一化输出的按列加权累加。整个过程全部在 VectorAIV流水上执行算子注册信息也印证了这一点Opcode::OP_ONLINE_SOFTMAX_UPDATE的核类型为 AIV、流水为 PIPE_V见 opcode.cpp。该接口通常与pypto.experimental.online_softmax配对使用online_softmax(scores, scale)对当前 scores 块做缩放并计算局部统计量exp 结果、列最大值、列指数和详见 online_softmax 文档online_softmax_update(...)把当前块统计量合入已有历史状态最终输出通常还需要用更新后的指数和做归一化即updated_output / updated_sum。函数原型online_softmax_update( previous_max: Tensor, previous_sum: Tensor, previous_output: Tensor, current_max: Tensor, current_sum: Tensor, current_output: Tensor, ) - Tuple[Tensor, Tensor, Tensor]Python 层的真实签名与文档一致见 operation.py其内部通过op_wrapper装饰后转发到 C 层pypto_impl.OnlineSoftmaxUpdate(...)构造OP_ONLINE_SOFTMAX_UPDATE算子节点。参数说明六个输入参数全部为二维 Tensor数据类型仅支持 DT_FP32参数名输入/输出说明previous_max输入历史块的列最大值。支持的数据类型为DT_FP32。不支持空 Tensor支持两维。Shape 为 [1, q_len]。previous_sum输入历史块的列指数和。支持的数据类型为DT_FP32。不支持空 Tensor支持两维。Shape 为 [1, q_len]。previous_output输入历史块累计的未归一化输出。支持的数据类型为DT_FP32。不支持空 Tensor支持两维。Shape 为 [head_dim, q_len]。current_max输入当前块的列最大值通常来自pypto.experimental.online_softmax。支持的数据类型为DT_FP32。Shape 为 [1, q_len]。current_sum输入当前块的列指数和通常来自pypto.experimental.online_softmax。支持的数据类型为DT_FP32。Shape 为 [1, q_len]。current_output输入当前块的未归一化输出。支持的数据类型为DT_FP32。Shape 为 [head_dim, q_len]需要与 previous_output 形状一致。其中统计量 Tensormax/sum在 operation_impl.cpp 的CheckOnlineSoftmaxUpdateStats中被逐项校验六个输入均须为 DT_FP32、两维、非空previousOutput与currentOutput形状必须一致四个 max/sum Tensor 形状必须互相一致max/sum 的 shape 必须是[1, q_len]且q_len与 output 的第二维相等即shape[0] 1 shape[1] previousOutput.shape[1]。返回值说明返回三个输出 Tensor全部为 DT_FP32返回值说明updated_max合并后的列最大值数据类型为 DT_FP32Shape 为 [1, q_len]。updated_sum合并后的列指数和数据类型为 DT_FP32Shape 为 [1, q_len]。updated_output合并后的未归一化输出数据类型为 DT_FP32Shape 为 [head_dim, q_len]。从实现看operation_impl.cpp三个输出的 Shape 分别继承自previousMax、previousSum、previousOutput同时算子还会申请一个额外的updateWorkspace形状为[head_dim 3, AlignUp(q_len, 32/4)]即[head_dim 3, 对齐到 8 列的 q_len]作为内部中间缓冲TileOp 内核把 workspace 按行切分为previousScaleTile、currentScaleTile、scaledCurrentSumTile、scaledCurrentOutputTile四块见 softmax.h。该 workspace 对用户不可见但解释了为什么约束中要求最后一维 Tile 需满足 FP32 的 32 字节对齐——GetOnlineSoftmaxFp32AlignedColumns正是按BLOCK_SIZE / sizeof(FP32)向上对齐列宽见 operation_impl.cpp。约束说明该接口为定制接口不保证稳定性。所有输入 Tensor 数据类型仅支持 DT_FP32。current_output 需要与 previous_output 形状一致。当前版本不切分第 0 维要求 previous_output.shape[0] vec_tile[0]。最后一条约束在源码中有精确的断言实现CheckOnlineSoftmaxTileShape会校验viewShape[0] vecTile[0]即 Tensor 的第 0 维必须能放进一个 Tile 的第 0 维operation_impl.cppCheckOnlineSoftmaxUpdateTileOperands还会进一步要求vecTile[1]能被BLOCK_SIZE / sizeof(DT_FP32)整除即最后一维 Tile 必须是 FP32 的 32 字节对齐operation_impl.cpp。此外统计量 Tensor 必须满足[1, q_len]的形状约束output 类 Tensor 形状必须两两一致违反任一条件都会触发ERR_PARAM_INVALID/ERR_CONFIG_TILE断言。调用示例TileShape 设置示例调用该 operation 接口前应通过set_vec_tile_shapes设置 TileShapeVector Tile 切分。TileShape 的维度设置须与previous_output、current_output保持一致当前版本不切分第 0 维要求previous_output.shape[0] vec_tile[0]最后一维 Tile 大小需要满足 FP32 的 32 字节对齐即vec_tile[1]为 8 的整数倍因为 8 个 FP32 恰好是 32 字节。接口调用示例import pypto previous_max pypto.tensor([1, 128], pypto.DT_FP32) previous_sum pypto.tensor([1, 128], pypto.DT_FP32) previous_output pypto.tensor([128, 128], pypto.DT_FP32) current_max pypto.tensor([1, 128], pypto.DT_FP32) current_sum pypto.tensor([1, 128], pypto.DT_FP32) current_output pypto.tensor([128, 128], pypto.DT_FP32) pypto.set_vec_tile_shapes(128, 64) updated_max, updated_sum, updated_output pypto.experimental.online_softmax_update( previous_max, previous_sum, previous_output, current_max, current_sum, current_output, )上述示例中set_vec_tile_shapes(128, 64)的第 0 维 128 满足previous_output.shape[0] 128 128第 1 维 64 是 8 的整数倍满足 FP32 32 字节对齐要求。仓库的 ST 冒烟测试给出了与此完全一致的、可编译运行的完整用例测试在pypto.function上下文中创建六个输入 Tensor设置pypto.set_vec_tile_shapes(128, 64)后调用该接口并断言updated_max/updated_sum形状为[1, 128]、updated_output形状为[128, 128]、数据类型均为DT_FP32见 test_online_softmax.py。典型应用场景FlashAttention 分块注意力中的逐块状态更新该接口最典型的落地场景是 FlashAttention 类 kernel在 KV 序列维按块迭代时每个 k-tile 用online_softmax计算局部统计量再通过online_softmax_update将新块统计量合入跨块累计状态。仓库中的 flash_attention_mha_impl.py 给出了完整参考实现Ascend 950 路径flash_attention_varlen_forward_950其核心循环结构如下首个 k-tilepij_bf16, mij, lij pypto.experimental.online_softmax(scores, scale)计算局部统计量并将mij、lij、oij写入累加器mi_update、li_update、oi_update后续 k-tile再次用online_softmax得到当前块统计量然后通过pypto.view取出累加器中的历史状态调用online_softmax_update(mi, li, oi, mij, lij, oij)得到mi_new, li_new, oi_tmp并把结果写回累加器flash_attention_mha_impl.py最后一个 k-tile用updated_sum归一化未归一化输出out_fp32 pypto.div(oi_tmp, li_new, ...)再转 BF16 写回输出flash_attention_mha_impl.py。注意online_softmax_update输出的是未归一化的合并输出最终结果必须用updated_sum做归一化这一点与原文档最终输出通常还需要用更新后的指数和做归一化的描述完全吻合。底层实现与代码路径速览层次文件说明Python 前端experimental/operation.pyonline_softmax_update的 Python 入口转发到pypto_impl.OnlineSoftmaxUpdate算子构造interface/operation/operation_impl.cpp参数校验、输出 Tensor 与 workspace 分配、算子节点创建Tile 切分interface/operation/operation_impl.cpp按vec_tile[1]沿第 1 维切分逐列生成TOnlineSoftmaxUpdateTileOp内核实现interface/tileop/vector/softmax.hTMAX/TSUB/TEXP/TMUL/TADD/TCOLEXPANDMUL完成在线 Softmax 合并形状推导interface/operation/op_infer_shape_impl.cpp输出 valid shape 继承自输入统计量与 output算子注册interface/operation/opcode.cppAIV 核、PIPE_V 流水TileOp 名为TOnlineSoftmaxUpdate代码生成codegen/npu/codegen_vector_unary_with_tmp.cpp生成带临时缓冲的完整参数 TileOp 调用测试用例tests/st/operation/vector/test_online_softmax.py950 SoC 冒烟测试验证 Shape 与 dtype其中值得注意的实现细节与online_softmax不同online_softmax_update没有标量属性其代码生成直接复用PrintTileOpWithFullParamsTmpBuf把 updateWorkspace 作为临时缓冲参与参数展开而online_softmax则需要把scale作为标量属性随算子下发见 codegen_vector_unary_with_tmp.cpp。总结pypto.experimental.online_softmax_update是一个面向 Ascend 950 系列、约束明确的在线 Softmax 状态合并算子。使用时应牢记三点全部输入输出均为 FP32 二维 Tensormax/sum 恒为[1, q_len]而 output 为[head_dim, q_len]且前后形状一致调用前必须通过set_vec_tile_shapes配置满足第 0 维不切分 最后一维 32 字节对齐的 TileShape。它与online_softmax组成局部统计 增量合并的完整在线 Softmax 流水配合pypto.div归一化即可支撑 FlashAttention 类分块注意力 kernel 的跨块数值稳定计算。由于该接口为定制接口且不保证稳定性接入业务前建议以当前版本仓库中的 ST 测试 和 FlashAttention 参考实现 为基线进行验证。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考