CANN pyasc 算子编程:asc.language.adv.softmax_flash_v2 在线 Softmax(FlashAttention-2)接口使用指南
2026/9/18 7:03:51 网站建设 项目流程

CANN pyasc 算子编程:asc.language.adv.softmax_flash_v2 在线 Softmax(FlashAttention-2)接口使用指南

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

本文围绕 CANN pyasc 项目中asc.language.adv.softmax_flash_v2高阶算子接口展开,系统讲解其 FlashAttention-2 对应的在线(online)Softmax 计算原理、完整参数语义、与 Ascend C 函数原型的对应关系、约束条件及多种调用形式,并给出基于仓库源码与真实用例(融合推理注意力示例)的可运行代码。读完本文,你将掌握如何在 pyasc 的 JIT 内核中正确声明状态张量、构造SoftmaxTiling/SoftMaxShapeInfo/SoftmaxConfig,并利用is_updateis_reuse_sourceis_basic_blockshared_tmp_buffer等参数完成首块与更新块两阶段的在线 Softmax 计算。

一、接口定位:SoftmaxFlash 增强版,对应 FlashAttention-2 算法

在自注意力(Self-Attention)计算中,QK 矩阵分数需要经过 Softmax 归一化后才能参与 PV 加权累加。传统实现需要一次性读出整行分数计算全局 max 与 sum,导致中间结果占用大量片上存储;而 FlashAttention 系列算法通过在线(online)Softmax技巧,将 max 与 sum 以流式方式维护,使得长序列场景下片上缓冲(UB)占用与序列长度解耦。

asc.language.adv.softmax_flash_v2是 pyasc 提供的 SoftmaxFlash 增强版本接口,官方文档明确其对应 FlashAttention-2 算法,用于在昇腾 AI 处理器上以分块(tile)方式完成带状态更新的 Softmax 计算。其对应的 Ascend C 函数原型为SoftmaxFlashV2,Python 接口与 Ascend C 接口一一对应,并遵守 Python 原生语法(参见 项目 README 与 python/asc/language/adv/activation.py 中的定义)。

从实现上看,该接口由 python/asc/language/adv/activation.py#L153-L160 中的softmax_flash_v2函数承载,函数体通过装饰器@require_jit标记,最终调用 IR 构建器的create_asc_SoftmaxFlashV2Op生成昇腾 IR 指令(见 activation.py#L341-L349),从而接入 pyasc 的 JIT 编译与执行链路。

二、函数签名与参数总览

接口定义位于 docs/python-api/language/generated/asc.language.adv.softmax_flash_v2.md,Python 侧签名为:

asc.language.adv.softmax_flash_v2( dst_tensor: LocalTensor, exp_sum_tensor: LocalTensor, max_tensor: LocalTensor, src_tensor: LocalTensor, exp_max_tensor: LocalTensor, in_exp_sum_tensor: LocalTensor, in_max_tensor: LocalTensor, tiling: SoftmaxTiling, softmax_shape_info: SoftMaxShapeInfo | None = None, shared_tmp_buffer: LocalTensor | None = None, out_reduce_max: LocalTensor | None = None, is_update: bool = False, is_reuse_source: bool = False, is_basic_block: bool = False, is_data_format_nz: bool = False, config: SoftmaxConfig | None = None ) -> None

参数语义可概括为"3 个状态输出、3 个状态输入 + 1 个数据输入"的结构:

参数方向说明
src_tensor源操作数待归一化的输入分数矩阵,last 轴长度需 32 Byte 对齐
dst_tensor目的操作数归一化后的概率输出,shape 与src_tensor一致
exp_sum_tensor目的操作数保存计算过程中 reducesum 的结果
max_tensor目的操作数保存计算过程中 reducemax 的结果
exp_max_tensor目的操作数保存in_max_tensor与本次 reducemax 差值的 e 指数幂结果
in_exp_sum_tensor源操作数上一分块传入的 sum 状态值
in_max_tensor源操作数上一分块传入的 max 状态值
tiling配置计算所需的SoftmaxTiling信息
softmax_shape_info配置(可选)src_tensor的 shape 信息,类型SoftMaxShapeInfo
shared_tmp_buffer临时空间(可选)数据类型固定为uint8,存储接口内部中间变量,由开发者提供
out_reduce_max目的操作数(可选)保存第一次 reducemax 的结果,shape 与max_tensor一致
is_update模板参数是否基于in_exp_sum_tensorin_max_tensor更新 softmax 状态,默认False
is_reuse_source模板参数是否复用src_tensor空间,默认False
is_basic_block模板参数是否使用基本块模式,默认False
is_data_format_nz模板参数输入输出是否为 NZ 格式,默认False
config模板参数(可选)SoftmaxConfig类型配置结构体

返回值:无(结果通过各目的操作数 Tensor 就地写出)。

三、与 Ascend C 函数原型的对应关系

pyasc 的 Python 接口与 Ascend C 函数一一对应。官方文档给出了SoftmaxFlashV2的 6 种重载形式,按"临时空间来源"和"数据类型是否一致"两个维度划分,理解这些原型有助于把握 Python 参数的映射关系。

3.1 按临时空间来源划分

接口框架自动申请临时空间(不传shared_tmp_buffer):

template <typename T, bool isUpdate = false, bool isReuseSource = false, bool isBasicBlock = false, bool isDataFormatNZ = false, const SoftmaxConfig& config = SOFTMAX_DEFAULT_CFG> __aicore__ inline void SoftmaxFlashV2( const LocalTensor<T>& dstTensor, const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& inExpSumTensor, const LocalTensor<T>& inMaxTensor, const SoftMaxTiling& tiling, const SoftMaxShapeInfo& softmaxShapeInfo = {})

通过sharedTmpBuffer入参传入临时空间(显式传入uint8_t类型临时缓冲,便于开发者复用已申请空间、降低重复分配开销):

template <typename T, bool isUpdate = false, bool isReuseSource = false, bool isBasicBlock = false, bool isDataFormatNZ = false, const SoftmaxConfig& config = SOFTMAX_DEFAULT_CFG> __aicore__ inline void SoftmaxFlashV2( const LocalTensor<T>& dstTensor, const LocalTensor<T>& outExpSum, const LocalTensor<T>& outMax, const LocalTensor<T>& srcTensor, const LocalTensor<T>& outExpMax, const LocalTensor<T>& inExpSum, const LocalTensor<T>& inMax, const LocalTensor<uint8_t>& sharedTmpBuffer, const SoftMaxTiling& tiling, const SoftMaxShapeInfo& softmaxShapeInfo = {})

3.2 按输出 ReduceMax 与数据类型是否一致划分

  • 输出 ReduceMax 的重载:在dstTensor之后插入outReduceMax参数,用于接收第一次 reducemax 的结果,即文档中的"LocalTensor 数据类型相同,且输出 ReduceMax"形式。
  • 数据类型不同的重载:当精度要求较高时,可使用half类型的dstTensor/srcTensor/expMaxTensor配合float类型的expSumTensor/maxTensor/inExpSumTensor/inMaxTensor,即"LocalTensor 数据类型不同,不输出 ReduceMax"形式(见 activation.py#L196-L208)。

Python 侧通过位置参数与关键字参数统一了这些重载:shared_tmp_bufferout_reduce_maxconfig均可选,框架会自动映射到对应 IR 操作(见 activation.py#L337-L349)。

四、核心参数详解与 32 Byte 对齐约束

4.1 状态张量的 last 轴 32 Byte 约束

在线 Softmax 需要跨分块传递 max 与 sum 状态。为保证向量单元以 datablock 为单位高效广播,文档对状态张量提出了统一的 last 轴约束:

  • exp_sum_tensor:保存 reducesum 结果。除配置为非拓展模式(SoftmaxMode.SOFTMAX_OUTPUT_WITHOUT_BRC)外,其 last 轴长度固定为 32 Byte,即一个 datablock 长度;该 datablock 中所有数据为同一个值(例如 float16 下该 block 内 16 个数均为相同的 reducesum 值)。非 last 轴长度与dst_tensor保持一致。
  • max_tensor:保存 reducemax 结果,last 轴同样固定为 32 Byte,block 内所有数据为同一个 reducemax 值,非 last 轴与dst_tensor一致。
  • exp_max_tensor:保存in_max_tensor与本次 reducemax 差值(in_max - max)的 e 指数幂结果,last 轴 32 Byte,block 内数据相同,非 last 轴与dst_tensor一致。
  • in_exp_sum_tensor:作为源操作数输入上一分块的 sum 值,last 轴 32 Byte,block 内数据相同。
  • in_max_tensor:作为源操作数输入上一分块的 max 值,last 轴 32 Byte,block 内数据相同。
  • src_tensor:源操作数,last 轴长度需要 32 Byte 对齐(注意这里是对齐而非固定 32 Byte)。

从仓库测试 python/test/unit/language/adv/test_activation.py#L52-L87 可以看出,实际使用中state_half被同时复用为exp_sum_tensormax_tensorin_exp_sum_tensorin_max_tensor,这正是文档"空间可复用"约束在实践中的典型形态。

4.2 tiling:SoftmaxTiling 结构

tiling参数类型为SoftmaxTiling,定义于 python/asc/language/adv/tiling.py#L88-L108。它是一个Struct,字段与昇腾asc_SoftMaxTilingType一一对应:

字段类型含义
src_m/src_k/src_sizeint32输入矩阵的 M、K 维度与元素总数(src_size = src_m * src_k
out_max_m/out_max_k/out_max_sizeint32max 状态输出的 M、K 维度与元素总数
split_m/split_k/split_sizeint32按 M、K 切分的分块维度与大小
reduce_m/reduce_k/reduce_sizeint32归约(reduce)维度信息
range_m/tail_mint32M 轴计算范围与尾块 M 维度
tail_split_size/tail_reduce_sizeint32尾块的分块与归约大小

构造示例:asc.adv.SoftmaxTiling(src_m=8, src_k=512, src_size=4096),即处理一个 8×512 的分数分块(8 行、512 列)。未指定的字段取默认值 0。

4.3 softmax_shape_info:SoftMaxShapeInfo 结构

softmax_shape_info类型为SoftMaxShapeInfo,定义于 python/asc/language/adv/types.py#L324-L343,用于描述逻辑维度与原始维度:

  • src_m/src_k:参与计算的逻辑矩阵维度;
  • ori_src_m/ori_src_k:GM 上原始输入矩阵的维度。

四者的组合支持补齐(padding)语义:当src_m != ori_src_msrc_k != ori_src_k时,需要将 GM 上的原始输入沿 M 轴或 K 轴补齐到逻辑维度,补齐数据会参与部分运算。若复用输入输出空间,结果会覆盖src_tensor中的补齐数据;否则覆盖dst_tensor中与补齐位置对应的数据。构造示例:asc.adv.SoftMaxShapeInfo(8, 512, 8, 512)(逻辑与原始维度一致,无补齐)。

4.4 config:SoftmaxConfig 与 SoftmaxMode

config为可选的编译期常量结构体SoftmaxConfig,定义于 python/asc/language/adv/types.py#L346-L375,构造参数为:

SoftmaxConfig(check_tiling: bool = True, src_m: int = 0, src_k: int = 0, mode: SoftmaxMode = SoftmaxMode.SOFTMAX_NORMAL)
  • check_tiling:是否在编译期检查 tiling 合法性,默认True
  • src_m/src_k:编译期已知的 shape 常量,用于 shape 特化;
  • mode:计算模式,取值来自 python/asc/language/core/enums.py#L215-L217 中的SoftmaxMode枚举:
    • SOFTMAX_NORMAL = 0:常规拓展模式,状态张量 last 轴按 32 Byte 拓展存储;
    • SOFTMAX_OUTPUT_WITHOUT_BRC = 1:非拓展模式,不进行 block 广播拓展。

模式与状态张量约束联动:除SOFTMAX_OUTPUT_WITHOUT_BRC场景外,exp_sum_tensormax_tensorexp_max_tensorin_exp_sum_tensorin_max_tensor的 last 轴长度必须固定为 32 Byte。SOFTMAX_OUTPUT_WITHOUT_BRC模式下状态张量无需按 32 Byte 拓展存储,适合输出后续不再依赖 block 广播的精简场景。

4.5 可选参数 out_reduce_max 的配套限制

out_reduce_max用于保存第一次 reducemax 的结果,shape 与max_tensor一致。指定该参数时需遵守以下限制:

  • is_updateFalse时,不输出该结果(即仅在更新阶段配合使用);
  • 仅支持 ND 格式,is_data_format_nz为预留参数,应使用默认值False
  • config.check_tiling为预留配置,应设为False
  • config.mode仅支持SoftmaxMode.SOFTMAX_OUTPUT_WITHOUT_BRC
  • 若将mode配置为SOFTMAX_NORMAL,接口不执行计算也不保存输出;
  • out_reduce_max外,其余输出的计算结果与未指定时相同。

4.6 布尔模板参数

参数默认值语义
is_updateFalse首块计算时置False(初始化 max/sum 状态),后续分块置True(基于in_exp_sum_tensor/in_max_tensor做在线更新)
is_reuse_sourceFalse是否复用src_tensor的空间保存输出,可降低片上存储占用
is_basic_blockFalse是否使用基本块模式(配合编译期 shape 特化,见下文实战)
is_data_format_nzFalse输入输出是否为 NZ 格式,默认 ND

五、使用约束(官方约束说明)

使用softmax_flash_v2前必须确认以下约束,否则计算结果或地址行为不符合预期:

  1. 空间复用规则src_tensordst_tensor的 Tensor 空间可以复用;max_tensorin_max_tensor的空间可以复用;exp_sum_tensorin_exp_sum_tensor的空间可以复用。测试用例中即用同一个state_half同时承担这 4 个状态参数。
  2. 32 Byte 约束:除SOFTMAX_OUTPUT_WITHOUT_BRC模式外,exp_sum_tensormax_tensorexp_max_tensorin_exp_sum_tensorin_max_tensor的 last 轴长度必须固定为 32 Byte。
  3. 地址对齐:操作数地址对齐要求遵循通用地址对齐约束。
  4. 临时空间隔离:不支持shared_tmp_buffer与源操作数或目的操作数地址重叠。
  5. 补齐语义src_m != ori_src_msrc_k != ori_src_k时,需在 GM 侧将原始输入补齐到逻辑维度,补齐数据会参与部分运算,且输出覆盖位置随是否复用输入输出而变化(详见 4.3 节)。

六、调用示例

6.1 官方文档最小示例(float16 全同型)

dst_half = asc.LocalTensor(dtype=asc.float16) state_half = asc.LocalTensor(dtype=asc.float16) src_half = asc.LocalTensor(dtype=asc.float16) exp_max_half = asc.LocalTensor(dtype=asc.float16) tiling = asc.adv.SoftmaxTiling(src_m=8, src_k=512, src_size=4096) shape = asc.adv.SoftMaxShapeInfo(8, 512, 8, 512) asc.adv.softmax_flash_v2( dst_half, state_half, state_half, src_half, exp_max_half, state_half, state_half, tiling, shape)

此处state_half同时作为exp_sum_tensormax_tensorin_exp_sum_tensorin_max_tensor,充分利用空间复用约束。

6.2 单元测试覆盖的完整组合形态

仓库单元测试 python/test/unit/language/adv/test_activation.py#L52-L87 展示了接口在 JIT 内核中的全部典型组合,可作为正确性参考:

@asc.jit def kernel_softmax_flash_v2() -> None: dst_half = asc.LocalTensor(dtype=asc.float16) state_half = asc.LocalTensor(dtype=asc.float16) src_half = asc.LocalTensor(dtype=asc.float16) exp_max_half = asc.LocalTensor(dtype=asc.float16) state_float = asc.LocalTensor(dtype=asc.float32) reduce_max = asc.LocalTensor(dtype=asc.float16) shared_tmp = asc.LocalTensor(dtype=asc.uint8) tiling = asc.adv.SoftmaxTiling(src_m=8, src_k=512, src_size=4096) shape = asc.adv.SoftMaxShapeInfo(8, 512, 8, 512) reduce_config = asc.adv.SoftmaxConfig(False, 8, 512, asc.SoftmaxMode.SOFTMAX_OUTPUT_WITHOUT_BRC) # ① 基础调用(float16 全同型,框架申请临时空间) asc.adv.softmax_flash_v2(dst_half, state_half, state_half, src_half, exp_max_half, state_half, state_half, tiling, shape) # ② 输出 ReduceMax + 在线更新(非拓展模式) asc.adv.softmax_flash_v2(dst_half, state_half, state_half, src_half, exp_max_half, state_half, state_half, tiling, shape, out_reduce_max=reduce_max, is_update=True, config=reduce_config) # ③ 状态张量使用 float32(数据类型不同形式:half 数据 + float 状态) asc.adv.softmax_flash_v2(dst_half, state_float, state_float, src_half, exp_max_half, state_float, state_float, tiling, shape) # ④ 显式传入 shared_tmp_buffer(uint8) asc.adv.softmax_flash_v2(dst_half, state_half, state_half, src_half, exp_max_half, state_half, state_half, tiling, shape, shared_tmp_buffer=shared_tmp) # ⑤ shared_tmp_buffer + out_reduce_max + is_update 组合 asc.adv.softmax_flash_v2(dst_half, state_half, state_half, src_half, exp_max_half, state_half, state_half, tiling, shape, shared_tmp_buffer=shared_tmp, out_reduce_max=reduce_max, is_update=True, config=reduce_config) # ⑥ 基本块模式(编译期全 tile 特化) full_tile_config = asc.adv.SoftmaxConfig(False, 8, 512) asc.adv.softmax_flash_v2(dst_half, state_float, state_float, src_half, exp_max_half, state_float, state_float, tiling, shape, shared_tmp_buffer=shared_tmp, is_basic_block=True, config=full_tile_config)

6.3 实战:融合推理注意力(Fused Infer Attention)中的两阶段在线 Softmax

仓库示例 examples/10_fused_infer_attention/fused_infer_attention.py 将softmax_flash_v2用于 FlashAttention-2 风格的融合注意力内核,是理解is_update语义的最佳实战参考。该示例中SoftmaxFlashV2每次处理 8 行分数(SOFTMAX_ROWS = 8KV_TILE = 512),以保持临时 UB 空间有界;QK 分数经 AIC 计算后交给 AIV 执行 Softmax,再以流水线方式与 PV 计算重叠(注释见 fused_infer_attention.py#L26-L53)。

两阶段调用逻辑封装在_online_softmax(fused_infer_attention.py#L163-L184):

@asc.jit def _online_softmax(prob_local, score_local, tile_sum, tile_max, tile_exp, shared_tmp, first_tiling, update_tiling, full_config, rows, kv_rows, inner): shape = asc.adv.SoftMaxShapeInfo(rows, kv_rows, rows, kv_rows) is_full_tile = rows == SOFTMAX_ROWS and kv_rows == KV_TILE if inner == 0: # 首块:is_update=False,初始化在线状态 if is_full_tile: asc.adv.softmax_flash_v2(prob_local, tile_sum, tile_max, score_local, tile_exp, tile_sum, tile_max, first_tiling, shape, shared_tmp_buffer=shared_tmp, is_basic_block=True, config=full_config) else: asc.adv.softmax_flash_v2(prob_local, tile_sum, tile_max, score_local, tile_exp, tile_sum, tile_max, first_tiling, shape, shared_tmp_buffer=shared_tmp, is_basic_block=True) else: # 后续分块:is_update=True,基于 in_exp_sum/in_max 更新状态 if is_full_tile: asc.adv.softmax_flash_v2(prob_local, tile_sum, tile_max, score_local, tile_exp, tile_sum, tile_max, update_tiling, shape, shared_tmp_buffer=shared_tmp, is_update=True, is_basic_block=True, config=full_config) else: asc.adv.softmax_flash_v2(prob_local, tile_sum, tile_max, score_local, tile_exp, tile_sum, tile_max, update_tiling, shape, shared_tmp_buffer=shared_tmp, is_update=True, is_basic_block=True)

该用例同时演示了三个实践要点:

  1. 首块/更新块语义inner == 0is_update=False初始化在线状态,inner > 0is_update=True做状态更新;tile_sumtile_max在调用间持续充当输入输出状态,实现跨 KV 分块的在线归一化。
  2. is_basic_block=True+full_config组合:当分块形状与编译期常量(8×512)完全匹配时(is_full_tile),传入SoftmaxConfig(False, SOFTMAX_ROWS, KV_TILE)进行 shape 特化,跳过 tiling 检查并生成更紧凑的基本块代码;形状不匹配时退回通用路径(仅is_basic_block=True,不传 config)。
  3. shared_tmp_buffer复用:显式传入uint8临时缓冲,避免每个分块重复申请临时空间。

对应的 tiling 构造逻辑参见build_softmax_tiling函数(fused_infer_attention.py#L751-L756),它按分数 tile 的形状返回SoftmaxTiling对象,供首块与更新块分别使用。

七、底层实现链路与验证方式

从源码结构可以梳理出该接口的完整实现链路:

  1. Python 接口层softmax_flash_v2定义于 python/asc/language/adv/activation.py#L153,入口位于asc.language.adv命名空间(由 python/asc/language/adv/init.py 导出)。
  2. IR 构建层:函数体将各参数经to_ir()转为 IR 句柄后调用create_asc_SoftmaxFlashV2Op,生成SoftmaxFlashV2昇腾方言算子(activation.py#L337-L349),算子定义见 include/ascir/Dialect/Asc/IR/Adv/Activation.td。
  3. 代码生成层lib/Target/AscendC/Adv/Activation.cpp将该算子翻译为 Ascend C 的SoftmaxFlashV2调用,Python 侧的函数原型注释即与生成的 C++ 模板签名一一对应(activation.py#L164-L254)。
  4. 配置结构体SoftmaxTiling(tiling.py#L88)、SoftMaxShapeInfoSoftmaxConfig(types.py#L324、types.py#L346)均为IRValue/Struct,在构造时即通过create_asc_ConstructOp生成编译期常量 IR。

验证方面,除上述单元测试外,还可参考:

  • MLIR 测试样例 test/Target/AscendC/adv.mlir,包含 Softmax 相关 op 的降级(lowering)结果,可用于检查 IR 形态;
  • 端到端测试 python/test/kernels/test_matmul.py、python/test/kernels/conftest.py 提供的 JIT 执行框架,配合asc.jit装饰器即可将内核编译并在 Model/昇腾后端运行。

八、小结

asc.language.adv.softmax_flash_v2是 pyasc 面向 FlashAttention-2 类算法提供的在线 Softmax 高阶接口,其核心价值在于以"3 输入 3 输出 + 状态张量 32 Byte 对齐"的约定,将跨分块的 max/sum 状态维护封装为可复用的算子调用。使用时抓住三个关键点即可正确落地:

  • 状态语义exp_sum/max/exp_max为输出,in_exp_sum/in_max为输入,首块is_update=False、后续分块is_update=True
  • 对齐约束:除SOFTMAX_OUTPUT_WITHOUT_BRC模式外,5 个状态张量 last 轴固定 32 Byte,src_tensorlast 轴 32 Byte 对齐;
  • 空间策略:按需组合shared_tmp_bufferis_reuse_sourceis_basic_block与编译期SoftmaxConfig特化,在 UB 占用与代码效率之间取得平衡。

如需进一步查看接口文档、算子定义与相关示例,可在仓库中继续阅读 docs/python-api/language/adv.md、docs/python-api/language/core.md(LocalTensor 说明)、docs/python-api/language/generated/asc.language.adv.softmax_flash_v2.md 以及 examples/10_fused_infer_attention/README.md。

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询