news 2026/9/19 0:08:48

PyPTO 在线 Softmax 状态更新算子 `online_softmax_update` 使用详解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyPTO 在线 Softmax 状态更新算子 `online_softmax_update` 使用详解

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、两维、非空;
  • previousOutputcurrentOutput形状必须一致;
  • 四个 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 分别继承自previousMaxpreviousSumpreviousOutput,同时算子还会申请一个额外的updateWorkspace(形状为[head_dim + 3, AlignUp(q_len, 32/4)],即[head_dim + 3, 对齐到 8 列的 q_len])作为内部中间缓冲:TileOp 内核把 workspace 按行切分为previousScaleTilecurrentScaleTilescaledCurrentSumTilescaledCurrentOutputTile四块(见 softmax.h)。该 workspace 对用户不可见,但解释了为什么约束中要求最后一维 Tile 需满足 FP32 的 32 字节对齐——GetOnlineSoftmaxFp32AlignedColumns正是按BLOCK_SIZE / sizeof(FP32)向上对齐列宽(见 operation_impl.cpp)。

约束说明

  1. 该接口为定制接口,不保证稳定性。
  2. 所有输入 Tensor 数据类型仅支持 DT_FP32。
  3. current_output 需要与 previous_output 形状一致。
  4. 当前版本不切分第 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_outputcurrent_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),其核心循环结构如下:

  1. 首个 k-tilepij_bf16, mij, lij = pypto.experimental.online_softmax(scores, scale)计算局部统计量,并将mijlijoij写入累加器mi_updateli_updateoi_update
  2. 后续 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);
  3. 最后一个 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.cppvec_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 二维 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 0:05:47

Navicat导出SQL脚本:表结构与数据导出的原理、避坑与实战

1. 项目概述&#xff1a;为什么导出SQL脚本是数据库日常工作的“保命操作”Navicat 是我用过最顺手的数据库可视化工具之一&#xff0c;不是因为它多炫酷&#xff0c;而是它把那些藏在命令行深处、容易手抖写错的 SQL 操作&#xff0c;变成了几个点击就能稳稳落地的动作。但很多…

作者头像 李华
网站建设 2026/9/19 0:05:42

DBSCAN聚类算法详解:从密度概念到Python实战与调参技巧

简介&#xff1a;这是题为《机器学习__DBSCAN算法》的PPT课件&#xff0c;面向机器学习初学者与数据挖掘实践者&#xff0c;系统讲解基于密度的聚类方法&#xff0c;解决传统K-Means需预设簇数、难以处理任意形状簇与噪声数据的问题。内容围绕核心点、边界点、噪声点三个关键概…

作者头像 李华
网站建设 2026/9/19 0:02:20

SYB创业计划书财务逻辑拆解:从销售收入预测到现金流量计划

简介&#xff1a;SYB创业计划书完整版.doc 是一份面向创业者、备赛学生及有开店打算人群的实用模板&#xff0c;以一家社区日用超市为案例&#xff0c;围绕企业概况、创业者个人情况、市场评估、市场营销计划、企业组织结构、固定资产、流动资金、销售收入预测、销售和成本计划…

作者头像 李华
网站建设 2026/9/18 23:56:05

SRDQN赋能多级供应链库存优化:从啤酒游戏到可部署决策

简介&#xff1a;本资源是一份面向科研人员与1–3年经验研发工程师的深度强化学习实践指南&#xff0c;聚焦供应链库存优化这一经典难题&#xff0c;以啤酒游戏为载体&#xff0c;系统复现并详解SRDQN算法在多级分散式供应链中的创新应用。资源直击牛鞭效应建模痛点&#xff0c…

作者头像 李华
网站建设 2026/9/18 23:53:34

Linux环境下IAR嵌入式工具链安装配置与命令行编译实践

很多嵌入式工程师一提IAR&#xff0c;脑子里第一反应就是Windows下的EWARM IDE。我自己干了这么多年固件开发&#xff0c;以前也是这个印象&#xff0c;直到公司开始搭CI流水线、要用Linux服务器统一出固件包&#xff0c;才不得不正视一个问题&#xff1a;IAR到底能不能在Linux…

作者头像 李华