CANN ops-transformer BlitzSparseAttention 算子实践:基于 sabi 块稀疏的 Prompt FlashAttention 全量推理优化
2026/9/20 12:26:51 网站建设 项目流程

CANN ops-transformer BlitzSparseAttention 算子实践:基于 sabi 块稀疏的 Prompt FlashAttention 全量推理优化

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

BlitzSparseAttention 是 CANN ops-transformer 实验分支中面向全量推理(prefill)场景的 FlashAttention 变体算子,它在 PromptFlashAttention 基础上引入sabi(Sparse Attention Block Index)块稀疏索引与block_shape块粒度控制,让 prefill 阶段的大语言模型(如 Hunyuan-video 等可受益于稀疏计算的端到端推理管线)只计算被选中的 KV 块,从而显著压缩注意力计算量。本文以 aclnnBlitzSparseAttention.md 为主体,结合仓库内算子定义、tiling 源码与 benchmark 脚本,完整讲解该算子的两段式 aclnn 接口、全部入参语义、约束条件、INT8 量化组合规则、C++/Python 调用方式与性能特征。读完本文,你将能够独立完成 BlitzSparseAttention 的调用参数配置、sabi 索引构造、样例编译运行与正确性/性能验证。

产品支持情况与算子功能定位

BlitzSparseAttention 在仓库中的产品支持情况如下:

  • Atlas A3 训练系列产品 / Atlas A3 推理系列产品:支持。
  • Atlas A2 训练系列产品 / Atlas A2 推理系列产品:支持。

该算子定位为全量推理场景的 FlashAttention 算子,核心特性包括:

  • sparse 优化:通过sabi指定每个 Q 块实际需要计算的 KV 块,跳过无关计算;
  • actualSeqLengthsKv 优化:key/value 的有效序列长度可与 query 解耦,支持 KV 侧变长;
  • INT8 量化:支持全量化(输入输出均为 INT8)与后量化(输入 FP16/BF16、输出 INT8)两条路径;
  • innerPrecise 参数:在高精度/高性能两种计算模式间选择,并支持行无效修正。

在算子定义文件 blitz_sparse_attention_def.cpp 中可以看到,该算子为 ascend910b(Atlas A2)与 ascend910_93(Atlas A3)注册了 AICore 配置,并额外为 ascend310p 注册了仅支持 FLOAT16 的精简配置(OpAICoreConfig config_310p),与"Atlas 推理系列产品仅支持 FLOAT16"的约束相互印证。

计算公式

自注意力(self-attention)利用输入样本自身的关联构建注意力模型:假设有一个长度为 $n$ 的输入样本序列 $x$,$x$ 的每个元素都是一个 $d$ 维向量,可将每个 $d$ 维向量视为一个 token embedding;序列经过 3 个权重矩阵变换得到 3 个维度为 $n \times d$ 的矩阵。$Q$、$K$、$V$ 即输入样本的重要属性元素,由输入样本经空间变换得到并统一到同一特征空间,公式及算子名称中的 "Attention" 为 "self-attention" 的简写:

$$ Attention(Q,K,V)=Score(Q,K)V $$

本算子中 Score 函数采用 Softmax 函数,self-attention 计算公式为:

$$ Attention(Q,K,V)=Softmax(\frac{QK^T}{\sqrt{d}})V $$

其中 $Q$ 与 $K^T$ 的乘积代表输入 $x$ 的注意力,为避免该值过大,通常除以 $d$ 的开根号进行缩放,再对每行做 softmax 归一化,最后与 $V$ 相乘得到 $n \times d$ 的矩阵。

核心机制:sabi 块稀疏索引与 block_shape 块粒度

sabi(Sparse Attention Block Index)是 BlitzSparseAttention 区别于普通 FlashAttention 的关键新增入参,其语义定义(见文档"约束说明"中 sabi 小节):

  • Shape[batch_size, num_heads, num_sabi_rows, num_sabi_cols],其中num_sabi_rows = ceil(sequence_length / BLOCK_SIZE_Q)num_sabi_cols = ceil(sequence_length / BLOCK_SIZE_KV)
  • 粒度来源BLOCK_SIZE_QBLOCK_SIZE_KV来自block_shape属性,两者默认均为 128,可取值集合为{128, 256, 512, 1024}
  • 数据类型uint16。每个元素是[0, num_sabi_cols)范围内的列索引,标识该 Q 行对应需要计算的BLOCK_SIZE_KV粒度 KV 块;未使用的槽位以0xFFFF(= 65535)填充,内核将其视为"跳过"标记。
  • 语义:对给定 batchb与 headhsabi[b, h, i, :]列出了第iBLOCK_SIZE_Q粒度 Q 行需要计算的 KV 块集合。

文档给出的示例(batch_size=1, num_heads=2, sequence_length=4096, block_shape=(128, 512),sabi 为 32 行 × 8 列):

[ # head 0: [ [0, 1, 2, 65535, 65535, 65535, 65535, 65535], # compute 3 of 8 KV chunks [0, 1, 2, 3, 65535, 65535, 65535, 65535], # compute 4 of 8 [0, 1, 2, 3, 4, 5, 6, 7], # compute all (dense row) # ... 32 rows total [0, 1, 2, 3, 4, 5, 6, 65535], # compute 7 of 8 ], # head 1: ... 32 rows total ]

从 benchmark 脚本 benchmark.py 的注释可以看到实现细节:BLOCK_SHAPES枚举了{128, 256, 512, 1024} × {128, 256, 512, 1024}全部 16 种块粒度组合,且内核的稀疏循环假设 sabi 每行中的保留块索引升序排列FirstGreaterEqual线性扫描 + chunkIdx 前向遍历),因此构造 sabi 时应保持索引有序。

函数原型:两段式 aclnn 接口

算子执行接口为两段式接口:必须先调用aclnnBlitzSparseAttentionGetWorkspaceSize获取计算所需 workspace 大小及包含算子计算流程的执行器,再调用aclnnBlitzSparseAttention执行计算。第一段接口完成入参校验与 workspace 规划,第二段接口在工作区上实际下发计算。

aclnnStatus aclnnBlitzSparseAttentionGetWorkspaceSize( const aclTensor *query, const aclTensor *key, const aclTensor *value, const aclTensor *pseShift, const aclTensor *attenMask, const aclTensor *sabi, const aclIntArray *actualSeqLengths, const aclIntArray *actualSeqLengthsKv, const aclTensor *deqScale1, const aclTensor *quantScale1, const aclTensor *deqScale2, const aclTensor *quantScale2, const aclTensor *quantOffset2, int64_t numHeads, double scaleValue, int64_t preTokens, int64_t nextTokens, char *inputLayout, int64_t numKeyValueHeads, int64_t sparseMode, int64_t innerPrecise, bool softmaxLseFlag, const aclIntArray *blockShape, const aclTensor *attentionOut, const aclTensor *softmaxLse, uint64_t *workspaceSize, aclOpExecutor **executor)
aclnnStatus aclnnBlitzSparseAttention( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream)

从封装源码 aclnn_blitz_sparse_attention.cpp 可见,两个入口函数分别转发到InnerBlitzSparseAttentionGetWorkspaceSizeInnerBlitzSparseAttention完成实际逻辑,接口头文件为 aclnn_blitz_sparse_attention.h。

aclnnBlitzSparseAttentionGetWorkspaceSize 参数详解

以下参数表完整继承自原文档(数据格式"ND"表示按维度顺序连续存储的普通张量,维度 3-4 对应 BSH/BNSD/BSND 等排布):

参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
query输入公式中的输入 Q保持与 key、value 的数据类型一致FLOAT16、BFLOAT16、INT8ND3-4×
key输入公式中的输入 K保持与 query、value 的数据类型一致FLOAT16、BFLOAT16、INT8ND3-4×
value输入公式中的输入 V保持与 query、key 的数据类型一致FLOAT16、BFLOAT16、INT8ND3-4×
pseShift输入位置编码不使用可传 nullptr;综合约束见约束说明FLOAT16、BFLOAT16ND4×
attenMask输入mask 矩阵不使用可传 nullptr;综合约束见约束说明BOOL、INT8、UINT8ND2-4×
sabi输入块稀疏索引矩阵不使用可传 nullptr;语义详见本文 sabi 小节UINT16ND4×
actualSeqLengths输入不同 Batch 中 query 的有效序列长度不指定可传 nullptr;综合约束见约束说明INT64TND1-
actualSeqLengthsKv输入不同 Batch 中 key/value 的有效序列长度不指定可传 nullptr;综合约束见约束说明INT64TND1-
deqScale1输入BMM1 后面的反量化因子支持 per-tensor;不使用可传 nullptrUINT64、FLOAT32ND1-
quantScale1输入BMM2 前面的量化因子支持 per-tensor;不使用可传 nullptrUINT64、FLOAT32ND1-
deqScale2输入BMM2 后面的反量化因子支持 per-tensor;不使用可传 nullptrUINT64、FLOAT32ND1-
quantScale2输入输出的量化因子支持 per-tensor、per-channel;不使用可传 nullptrUINT64、FLOAT32ND1-
quantOffset2输入输出的量化偏移支持 per-tensor、per-channel;不使用可传 nullptrFLOAT32ND1-
numHeads输入query 的 head 个数-INT64ND1-
scaleValue输入公式中 d 开根号的倒数数据类型与 query 满足数据类型推导规则;用户不特意指定时建议传入 1.0DOUBLE-1-
preTokens输入attention 需要和前几个 Token 计算关联不特意指定时建议传入 2147483647INT64-1-
nextTokens输入attention 需要和后几个 Token 计算关联不特意指定时建议传入 0INT64-1-
inputLayout输入标识输入 query、key、value 的数据排布格式不特意指定时建议传入 "BSH";综合约束见约束说明CHAR---
numKeyValueHeads输入key、value 中 head 个数不特意指定时建议传入 0(表示与 query 相等);综合约束见约束说明INT64---
sparseMode输入sparse 的模式综合约束见约束说明INT64-1-
innerPrecise输入高精度或者高性能选择综合约束见约束说明INT8-1-
attentionOut输出公式中的输出-FLOAT16、BFLOAT16、INT8ND3-4-
workspaceSize输出返回用户需要在 Device 侧申请的 workspace 大小---1-
executor输出返回 op 执行器,包含算子计算流程---1-

返回值:返回aclnnStatus状态码,具体参见 aclnn返回码。第一段接口完成入参校验,若出现以下错误码,对应原因为:

返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或必选属性且为空指针时返回。
ACLNN_ERR_PARAM_INVALID161002query、key、value、pseShift、attenMask、attentionOut 的数据类型和数据格式不在支持范围内。
ACLNN_ERR_RUNTIME_ERROR361001API 内部调用 npu runtime 的接口异常。

aclnnBlitzSparseAttention 第二段接口

第二段接口负责在工作区上执行计算,参数如下:

参数名输入/输出描述
workspace输入在 Device 侧申请的 workspace 内存地址。
workspaceSize输入在 Device 侧申请的 workspace 大小,由第一段接口获取。
executor输入op 执行器,包含算子计算流程。
stream输入指定执行任务的 AscendCL stream 流。

同样返回aclnnStatus状态码,具体参见 aclnn返回码。注意两段式接口的调用约束:第二段接口不可重复调用,且 workspace 是除输入/输出外算子在 NPU 上完成计算所需的临时内存(参见 two_phase_api.md)。

约束说明(完整参数限制清单)

通用约束

  • 该接口与 PyTorch 配合使用时,需要保证 CANN 相关包与 PyTorch 相关包的版本匹配。
  • 入参为空的处理:算子内部判断参数 query 是否为空,为空则直接返回;query 为非空 Tensor、key/value 为空 Tensor(即 S2 为 0)时,attentionOut 填充为全零;attentionOut 为空 Tensor 时由 AscendCLNN 框架处理;其余标注了"可传入 nullptr"的入参为空指针时不进行处理。
  • 维度符号约定:B(Batch)为输入样本批量大小,S(Seq-Length)为序列长度,H(Head-Size)为隐藏层大小,N(Head-Num)为多头数,D(Head-Dim)为隐藏层最小单元尺寸且满足 D=H/N,T 表示所有 Batch 输入样本序列长度的累加和。

query/key/value 输入限制(Atlas A2/A3)

  • B 轴支持小于等于 65536(64k);当输入类型包含 INT8 且 D 轴非 32 对齐,或输入类型为 FLOAT16/BFLOAT16 且 D 轴非 16 对齐时,B 轴仅支持到 128。
  • N 轴支持小于等于 256。
  • S 支持小于等于 20971520(20M)。部分长序列场景下,计算量过大可能导致 bsa 算子执行超时(aicore error 类型报错,errorStr 为timeout or trap error),此时建议做 S 切分处理。计算量受 B、S、N、D 影响,值越大计算量越大。典型的会超时(B、S、N、D 乘积较大)的场景包括但不限于:
BQ_NQ_SDKV_NKV_S
120209715225612097152
1220971520256220971520
201209715225612097152
110209715251212097152
  • D 轴支持小于等于 512;inputLayout 为 BSH 或 BSND 时,要求 N*D 小于 65535。
  • 数据类型支持 FLOAT16、BFLOAT16、INT8。

TND 场景下 query/key/value 的综合限制(Atlas A2/A3)

  • T 小于等于 65536。
  • N 等于 8/16/32/64/128,且 Q_N、K_N、V_N 相等。
  • Q_D、K_D 等于 192,V_D 等于 128/192。
  • 数据类型仅支持 BFLOAT16。
  • sparse 模式仅支持 sparse=0 且不传 mask,或 sparse=3 且传入 mask。
  • sparse=3 时,要求每个 batch 单独的 actualSeqLengths < actualSeqLengthsKv。

pseShift 使用限制

  • 输入 shape 需为(B, N, Q_S, KV_S)(1, N, Q_S, KV_S),其中 Q_S 为 query 的 shape 中的 S,KV_S 为 key 和 value 的 shape 中的 S。
  • Q_S 需大于等于 query 的 S 长度,KV_S 需大于等于 key 的 S 长度。
  • pseShift 的 KV_S 非 32 对齐时,建议 padding 到 32 字节提升性能,多余部分填充值不做要求;不使用时可传 nullptr。
  • Atlas A2/A3:数据类型支持 FLOAT16、BFLOAT16;pseShift 为 FLOAT16 时要求 query 为 FLOAT16 或 INT8,pseShift 为 BFLOAT16 时要求 query 为 BFLOAT16。query/key/value 为 FLOAT16 且 pseShift 存在时,默认走高精度模式,对应限制继承自高精度模式。

attenMask 使用限制

  • attenMask 的 KV_S 非 32 对齐时,建议 padding 到 32 对齐提升性能,多余部分填充成 1。
  • 通常建议 shape 为Q_S, KV_SB, Q_S, KV_S1, Q_S, KV_SB, 1, Q_S, KV_S1, 1, Q_S, KV_S,其中 Q_S 为 query 的 S,KV_S 为 key/value 的 S。
  • Atlas A2/A3:数据类型支持 BOOL、INT8 和 UINT8。

actualSeqLengths / actualSeqLengthsKv 使用限制

  • actualSeqLengths 不指定时可传 nullptr,表示有效序列长度与 query 的 shape 中 S 长度相同;注意每个 batch 的有效序列长度不应超过 query 对应 batch 的序列长度。
  • actualSeqLengthsKv 不指定时可传 nullptr,表示有效序列长度与 key/value 的 shape 中 S 长度相同;每个 batch 的有效序列长度不应超过 key/value 对应 batch 的序列长度。
  • 传入长度规则:长度为 1 时所有 Batch 使用相同 seqlen;长度大于等于 Batch 数量时取前 Batch 个数值;其他长度不被支持。
  • 当 query 的 inputLayout 为 TND 时,该入参必须传入,且以该入参元素数量作为 Batch 值;每个元素值表示当前 Batch 与之前所有 Batch 的 Sequence Length 和,因此后一个元素值必须大于等于前一个元素值,且不能出现负值。
  • Atlas A2/A3:数据类型支持 INT64,支持 TND 格式。

量化因子与 preTokens/nextTokens 限制(Atlas A2/A3)

  • deqScale1、deqScale2:数据类型支持 UINT64、FLOAT32。
  • quantScale1:数据类型支持 FLOAT32。
  • quantScale2、quantOffset2:数据类型支持 FLOAT32 和 BFLOAT16。
  • preTokens、nextTokens:数据类型支持 INT64。

inputLayout 使用限制

  • 当前支持 BSH、BSND、BNSD、BNSD_BSND(输入为 BNSD 时,输出格式为 BSND);用户不特意指定时建议传入 "BSH"。
  • Atlas A2/A3 除上述格式外还支持 TND(TND 不支持 pse、全量化、后量化)。

numKeyValueHeads 使用限制

  • 需要满足 numHeads 整除 numKeyValueHeads,且两者比值不能大于 64;在 BSND、BNSD、BNSD_BSND 场景下需与 shape 中 key/value 的 N 轴 shape 值相同,否则报错。
  • Atlas A2/A3:数据类型支持 INT64。

sparseMode 使用限制(Atlas A2/A3)

  • sparseMode=0(defaultMask):未传入 attenmask 时不执行 mask 操作,忽略 preTokens 和 nextTokens(内部赋值为 INT_MAX);传入时需传入完整 attenmask 矩阵(S1*S2),表示 preTokens 和 nextTokens 之间部分需要计算。
  • sparseMode=1(allMask):必须传入完整 attenmask 矩阵(S1*S2)。
  • sparseMode=2(leftUpCausal):传入优化后的 attenmask 矩阵(2048*2048)。
  • sparseMode=3(rightDownCausal):以右顶点为划分的下三角场景,传入优化后的 attenmask 矩阵(2048*2048)。
  • sparseMode=4(band):传入优化后的 attenmask 矩阵(2048*2048)。
  • sparseMode=5、6、7、8:分别代表 prefix、global、dilated、block_local,均暂不支持;用户不特意指定时建议传入 0。
  • 当 inputLayout 为 TND 时,sparseMode 仅支持取值 0、3。

关于各 sparseMode 的通用语义背景(defaultMask、allMask、leftUpCausal、rightDownCausal、band、prefix、global、dilated、block_local、treeMask 等模式的说明),可参见 sparse_mode_introduction.md。

innerPrecise 使用限制

共 4 种模式,共两位 bit 位:第 0 位(bit0)表示高精度或高性能选择,第 1 位(bit1)表示是否做行无效修正:

innerPrecise模式行无效修正
0高精度×
1高性能×
2高精度
3高性能
  • Q_S>1 时,sparse_mode 为 0 或 1 并传入用户自定义 mask 的情况下,建议开启行无效。
  • BFLOAT16 和 INT8 不区分高精度和高性能,行无效修正对 FLOAT16、BFLOAT16 和 INT8 均生效。当前 0、1 为保留配置值,当计算过程中"参与计算的 mask 部分"存在某整行全为 1 时精度可能有损失,可将该参数配置为 2 或 3 开启行无效提升精度,但会导致性能下降。若算子可判断出存在无效行场景(例如 sparse_mode 为 3、Sq > Skv),会自动开启无效行计算。

attentionOut 输出限制

  • 当 inputLayout 为 BNSD_BSND 时,输入 query 的 shape 是 BNSD,输出 shape 为 BSND;其余情况 attentionOut 的 shape 需与 query 保持一致。
  • Atlas A2/A3:数据类型支持 FLOAT16、BFLOAT16、INT8。

INT8 量化相关入参与输入/输出格式的综合限制

  • 输入 INT8、输出 INT8:deqScale1、quantScale1、deqScale2、quantScale2 需同时存在,quantOffset2 可选,不传时默认为 0。
  • 输入 INT8、输出 FLOAT16:deqScale1、quantScale1、deqScale2 需同时存在;若存在 quantOffset2 或 quantScale2(不为 nullptr)则报错并返回。
  • 输入 FLOAT16/BFLOAT16、输出 INT8:quantScale2 需存在,quantOffset2 可选(默认为 0);若存在 deqScale1、quantScale1 或 deqScale2(不为 nullptr)则报错并返回。
  • quantScale2 和 quantOffset2 支持 per-tensor/per-channel 两种格式与 FLOAT32/BFLOAT16 两种数据类型;传入 quantOffset2 时需保证其类型和 shape 信息与 quantScale2 一致。输入为 BFLOAT16 时同时支持 FLOAT32 和 BFLOAT16,否则仅支持 FLOAT32。per-channel 格式下,输出 layout 为 BSH 时要求 quantScale2 所有维度的乘积等于 H,其他 layout 要求乘积等于 N*D(建议:输出 layout 为 BSH 时 quantScale2 shape 传[1,1,H][H];输出为 BNSD 时传[1,N,1,D][N,D];输出为 BSND 时传[1,1,N,D][N,D])。
  • 输出为 INT8 且 quantScale2、quantOffset2 为 per-channel 时,暂不支持左 padding、Ring Attention 或 D 非 32 Byte 对齐场景。
  • 输出为 INT8 时,暂不支持 sparse 为 band 且 preTokens/nextTokens 为负数。

输出 INT8 + quantOffset2 的拦截场景

当输出为 INT8、quantOffset2 传入非空指针和非空 tensor 值,且 sparseMode、preTokens、nextTokens 满足以下条件时,矩阵存在某几行不参与计算,会导致计算结果误差,该场景会被拦截(解决方案:如需该场景不被拦截,需在 BSA 接口外部做后量化操作,不在 BSA 接口内部开启):

  • sparseMode = 0:attenMask 非空指针时,每个 batch 满足actualSeqLengths - actualSeqLengthsKV - preTokens > 0nextTokens < 0即满足拦截条件。
  • sparseMode = 1 或 2:不会出现满足拦截条件的情况。
  • sparseMode = 3:每个 batch 满足actualSeqLengthsKV - actualSeqLengths < 0即满足拦截条件。
  • sparseMode = 4:preTokens < 0或每个 batch 满足nextTokens + actualSeqLengthsKV - actualSeqLengths < 0即满足拦截条件。

调用示例:C++ aclnn 接口完整流程

示例代码(完整版见 test_aclnn_blitz_sparse_attention.cpp)演示了从 ACL 初始化、tensor 构造、两段式调用到结果回读与资源释放的完整流程。核心执行函数如下:

int ExecuteBlitzSparseAttention(TensorResources& resources, aclrtStream stream, void** workspaceAddr, uint64_t* workspaceSize) { int64_t numHeads = 8; int64_t numKeyValueHeads = 8; float scaleValue = static_cast<float>(1.0 / sqrt(128.0)); int64_t preTokens = 65535; // 覆盖所有前序 token int64_t nextTokens = 65535; // 覆盖所有后序 token int64_t sparseMode = 0; // defaultMask 模式 int64_t innerPrecise = 1; // 高性能模式 bool softmaxLseFlag = true; // 请求输出 softmax_lse // block_shape 属性:[BLOCK_SIZE_Q, BLOCK_SIZE_KV] std::vector<int64_t> blockShapeVec = {128, 128}; aclIntArray* blockShape = aclCreateIntArray(blockShapeVec.data(), blockShapeVec.size()); constexpr const char LAYER_OUT_STR[] = "BNSD"; char layerOut[sizeof(LAYER_OUT_STR)]; memcpy(layerOut, LAYER_OUT_STR, sizeof(LAYER_OUT_STR)); aclOpExecutor* executor; // 第一段:获取 workspace 大小与执行器 int ret = aclnnBlitzSparseAttentionGetWorkspaceSize( resources.queryTensor, resources.keyTensor, resources.valueTensor, nullptr /*pseShift*/, nullptr /*attenMask*/, resources.sabiTensor, resources.actualSeqLengths, resources.actualSeqLengthsKv, nullptr /*deqScale1*/, nullptr /*quantScale1*/, nullptr /*deqScale2*/, nullptr /*quantScale2*/, nullptr /*quantOffset2*/, numHeads, scaleValue, preTokens, nextTokens, layerOut, numKeyValueHeads, sparseMode, innerPrecise, softmaxLseFlag, blockShape, resources.outTensor, resources.lseTensor, workspaceSize, &executor); if (!CHECK_RET(ret == ACL_SUCCESS)) { /* 错误处理 */ } // 按 workspaceSize 在 Device 侧申请内存 if (*workspaceSize > 0ULL) { ret = aclrtMalloc(workspaceAddr, *workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); if (!CHECK_RET(ret == ACL_SUCCESS)) { /* 错误处理 */ } } // 第二段:执行计算 ret = aclnnBlitzSparseAttention(*workspaceAddr, *workspaceSize, executor, stream); if (!CHECK_RET(ret == ACL_SUCCESS)) { /* 错误处理 */ } aclDestroyIntArray(blockShape); return ACL_SUCCESS; }

示例中的 tensor 构造要点:

  • shape 设计:query/key/value 均为[B, N, S, D] = [1, 8, 512, 128](BNSD 排布),sabi 为[1, 8, 4, 4]S/128=4个 Q 块 × 4 个 KV 块),LSE 为[1, 8, 512](FLOAT32)。示例注释特别提示:需保证N*S/128 > coreNum(约 20),使 tiling 保持sOuterFactor=128(BSA 内核唯一编译的 tile 尺寸),否则 tiling 退化到sOuterFactor=32,找不到匹配的编译模板,会触发ADD_TO_LAUNCHER_LIST_AICORE失败并报错 561103。
  • sabi 数据:密集场景下每行填满保留的 KV 块索引0..KV_TILES-1;稀疏场景下用0xFFFF填充未使用槽位。
  • 全 1 输入的期望值校验:对全 1 输入,score =D*(1/sqrt(D)) = sqrt(D),LSE =log(S * exp(sqrt(D))) = log(S) + sqrt(D),示例中 S=512、D=128 时期望 LSE ≈ 17.5520。
  • 主流程Init(aclInit / aclrtSetDevice / aclrtCreateStream)→InitializeTensors(构造各 aclTensor 与 aclIntArray)→ExecuteBlitzSparseAttentionaclrtSynchronizeStreamProcessResults(Device 到 Host 拷贝并打印)→CleanupResources(销毁 tensor、释放内存、销毁 stream、aclFinalize)。

原文档附带的另一版调用示例位于 aclnnBlitzSparseAttention.md 的"调用示例"小节,其使用 BNSD 排布、sparseMode=0innerPrecise=1numHeads=2scaleValue=1/sqrt(2)preTokens=nextTokens=65535的 FLOAT16 场景,展示了相同接口在不同 shape 下的复用方式。

编译与运行

示例编译与运行遵循仓库通用流程(详见 compile_and_run_sample.md),README 提供了两条命令:

# 构建算子自定义实验包并安装,然后安装 torch_bsa torch 接口 bash build.sh --make_clean --experimental -j96 --pkg --soc=ascend910b --ops=blitz_sparse_attention ./build/cann-ops-transformer-custom_linux-"$(uname -i)".run (cd experimental/attention/blitz_sparse_attention/torch_interface && bash build.sh custom) # 运行纯 C++ 样例 bash build.sh --experimental --run_example blitz_sparse_attention eager cust --soc=ascend910b --vendor_name=custom

样例输出片段(Ascend910B,全 1 输入):

query: [1, 8, 512, 128] (B, N, S, D) fp16 key: [1, 8, 512, 128] (B, N, S, D) fp16 value: [1, 8, 512, 128] (B, N, S, D) fp16 sabi: [1, 8, 4, 4] (B, N, Q_tiles, KV_tiles) uint16 out: [1, 8, 512, 128] (B, N, S, D) fp16 lse: [1, 8, 512] (B, N, S) float32 output[0] = 0x3C00 lse[0] = 17.550825

Python 侧快速验证:torch_bsa 接口

仓库为快速端到端验证提供了torch_bsa的 torch 接口(torch_interface.cpp 与 ascendc_extension.py),调用方式兼容torch_npu约定:

import torch import torch_bsa # sabi 粒度:两个值均需在 {128, 256, 512, 1024} 中;值越小块级控制越精细,但 sabi 张量越大。 # 默认(省略 block_shape 时)为 [128, 128]。 BLOCK_SIZE_Q, BLOCK_SIZE_KV = 128, 128 # sabi: torch.uint16,shape [B, N, ceil(S/BLOCK_SIZE_Q), ceil(S/BLOCK_SIZE_KV)]。 # 每行列出该 Q 块保留的 KV 块列索引,右侧以 0xFFFF(uint16 "skip" 哨兵)填充。 sabi = ... # 根据你的稀疏模式构建 # 返回 (attention_out, softmax_lse)。 # softmax_lse 为 [B, N, S] float32(softmax_lse_flag=True 时), # 或 {0} shape 空 tensor(softmax_lse_flag=False,默认)。 attention_out, softmax_lse = torch_bsa.blitz_sparse_attention( q, k, v, sabi=sabi, actual_seq_lengths=actseqlen, actual_seq_lengths_kv=actseqlenkv, num_heads=h, num_key_value_heads=h, input_layout='BNSD', scale_value=scale, sparse_mode=0, softmax_lse_flag=False, # 置 True 可同时返回 log-sum-exp 输出 block_shape=[BLOCK_SIZE_Q, BLOCK_SIZE_KV], )

softmax_lse 输出说明

属性
控制参数softmax_lse_flag(bool 属性,默认 False)
输出位置输出索引 1(始终返回;flag 为 False 时为空)
使能时 shape[B, N, S]
数据类型float32(与 Q/K/V 数据类型无关)
Layout仅非 TND layout(BNSD、BSH、BSND);TND 返回{0}
语义每个 query 的 log-sum-exp:log(Σ exp(q·kᵀ / √d)),对所有被注意到的 KV token 求和

LSE 与注意力输出在同一次 kernel pass 内计算完成,不产生额外内存带宽开销;可用于 ring attention、投机解码重缩放(speculative decoding rescaling)等需要跨段合并 partial attention 结果的场景。当softmax_lse_flag=False时,kernel 跳过 LSE 写出路径并返回零元素占位 tensor,调用方无需为其分配内存。

性能基准与当前已知限制

benchmark 目录(benchmark.py、plot.py、test_attn.py、test_lse.py、test_joint.py)提供了一套完整的验证手段:

cd experimental/attention/blitz_sparse_attention/benchmark pytest test_attn.py # attention_out 正确性:序列长度 10k-30k、1-4 个注意力头,与 npu_fusion_attention 及自实现对比 pytest test_lse.py # softmax_lse 正确性:与 npu_fused_infer_attention_score 对比 pytest test_joint.py # 同时校验两个输出 python benchmark.py # 性能基准:对 BLOCK_SHAPES 中每对粒度扫描各稀疏度

README 中给出的 Ascend910B2 上的部分测量结果(BFLOAT16、BNSD layout、S=118806、D=128、B=1)显示:128×256 粒度在稀疏度约 0.05 时与稠密 PFA 参考实现打平,128×512 约在 0.1 时打平;历史记录的稀疏度 0.5 时 1.89× 加速对 128×512 依然成立,而 128×256 在保持 2 倍更细 sabi 分辨率的同时距离该加速比约 10%。完整表格与逐稀疏度 PFA 加速比汇总见 benchmark/README.md。

README 同时披露了当前实验状态的已知限制,使用前务必注意:

  • TODO 1(128×128 粒度加速有限):128×512 粒度在 50% 稀疏度下已表现出明显加速(1.89×),但当前 128×128 版本内部仍使用 128×512 matmul tile,未先将 sabi 选中的 128×128 子块压缩进 tile,导致 cube 会执行部分冗余 matmul 后再在 softmax 阶段 mask 掉——只有整段 512 长 tile 全空(即连续 4 个未选中块)才会被真正跳过,因此加速从稀疏度 ≥10% 开始按连续未选中块的概率爬升。仓库注释指出,彻底修复需要重写 matmul 调度;同门算子attention/block_sparse_attention采用自底向上的 CATLASS 设计,对这一问题处理得更好。
  • TODO 2(B>1 不正确):目前仅B=1已知结果正确,B>1会产出错误输出。所有测试与 benchmark 在问题修复前都必须以B=1运行。

源码实现佐证:算子注册、tiling 与封装

  • 算子定义:blitz_sparse_attention_def.cpp 定义了全部 13 个输入(query/key/value/pse_shift/atten_mask/sabi/actual_seq_lengths/actual_seq_lengths_kv/deq_scale1/quant_scale1/deq_scale2/quant_scale2/quant_offset2)、2 个输出(attention_out/softmax_lse)与 8 个属性(num_heads、scale_value、pre_tokens、next_tokens、input_layout、num_key_value_heads、sparse_mode、inner_precise、softmax_lse_flag、block_shape),其中softmax_lse_flag默认 False、block_shape默认{128, 128}pre_tokens默认 214748647、inner_precise默认 1。注册的 AICore 配置包括 ascend910b、ascend910_93 与 ascend310p。
  • tiling 常量:blitz_sparse_attention_tiling_const.h 中可见HIGH_PRECISION=0HIGH_PERFORMANCE=1MAX_BATCH=256等常量,以及 load 均衡系数数组COF[8] = {256, 384, 512, 640, 768, 896, 960, 1024},印证了 innerPrecise 的双模式语义与 batch 上限设计。
  • aclnn 封装:aclnn_blitz_sparse_attention.cpp 实现两个对外入口并转发至 inner 实现,头文件 aclnn_blitz_sparse_attention.h 声明接口。
  • kernel 侧:blitz_sparse_attention.cpp 及其 base 头文件构成算子核实现,sabi 的"跳过"语义(0xFFFF哨兵)在核侧被解析为跳过分块计算。

小结

BlitzSparseAttention 为全量推理场景提供了"以 sabi 块索引驱动、以 block_shape 控制粒度"的块稀疏 FlashAttention 实现,同时保留 actualSeqLengthsKv 变长优化、INT8 全量化/后量化与 innerPrecise 精度模式选择。使用时的关键动作可归纳为:按block_shape粒度构造升序排列、0xFFFF填充的 sabi 索引;按 sparseMode 语义选择 mask 传入方式;按量化组合规则配对 scale/offset 入参;以两段式 aclnn 接口申请 workspace 后执行;并在实验阶段严格遵守 B=1 与块粒度相关的已知限制。结合 benchmark 测试套件 可完成正确性与性能的闭环验证。

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

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

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

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

立即咨询