PyPTO 在线 Softmax 状态更新算子online_softmax_update使用详解
【免费下载链接】pyptoPyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
导读
pypto.experimental.online_softmax_update是 CANN PyPTO(Parallel 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.cpp),ST 冒烟测试也统一打上了pytest.mark.soc("950")标记,仅面向 950 系列 SoC 运行(见 test_online_softmax.py)。
功能说明:在线 Softmax 的状态合并
在线(online)Softmax 是长序列分块注意力中避免为整行 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,再用TSUB+TEXP分别求出历史块与当前块的缩放因子exp(previousMax - updatedMax)、exp(currentMax - updatedMax),随后TMUL+TADD合并指数和,TCOLEXPANDMUL+TADD完成对未归一化输出的按列加权累加。整个过程全部在 Vector(AIV)流水上执行,算子注册信息也印证了这一点: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 形状一致。 |
其中统计量 Tensor(max/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_FP32,Shape 为 [1, q_len]。 |
| updated_sum | 合并后的列指数和,数据类型为 DT_FP32,Shape 为 [1, q_len]。 |
| updated_output | 合并后的未归一化输出,数据类型为 DT_FP32,Shape 为 [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.cpp);CheckOnlineSoftmaxUpdateTileOperands还会进一步要求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设置 TileShape(Vector 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-tile:
pij_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.py | online_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.h | TMAX/TSUB/TEXP/TMUL/TADD/TCOLEXPANDMUL完成在线 Softmax 合并 |
| 形状推导 | interface/operation/op_infer_shape_impl.cpp | 输出 valid shape 继承自输入统计量与 output |
| 算子注册 | interface/operation/opcode.cpp | AIV 核、PIPE_V 流水,TileOp 名为TOnlineSoftmaxUpdate |
| 代码生成 | codegen/npu/codegen_vector_unary_with_tmp.cpp | 生成带临时缓冲的完整参数 TileOp 调用 |
| 测试用例 | tests/st/operation/vector/test_online_softmax.py | 950 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 二维 Tensor;max/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-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考