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_Q与BLOCK_SIZE_KV来自block_shape属性,两者默认均为 128,可取值集合为{128, 256, 512, 1024}。 - 数据类型:
uint16。每个元素是[0, num_sabi_cols)范围内的列索引,标识该 Q 行对应需要计算的BLOCK_SIZE_KV粒度 KV 块;未使用的槽位以0xFFFF(= 65535)填充,内核将其视为"跳过"标记。 - 语义:对给定 batch
b与 headh,sabi[b, h, i, :]列出了第i个BLOCK_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 可见,两个入口函数分别转发到InnerBlitzSparseAttentionGetWorkspaceSize与InnerBlitzSparseAttention完成实际逻辑,接口头文件为 aclnn_blitz_sparse_attention.h。
aclnnBlitzSparseAttentionGetWorkspaceSize 参数详解
以下参数表完整继承自原文档(数据格式"ND"表示按维度顺序连续存储的普通张量,维度 3-4 对应 BSH/BNSD/BSND 等排布):
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| query | 输入 | 公式中的输入 Q | 保持与 key、value 的数据类型一致 | FLOAT16、BFLOAT16、INT8 | ND | 3-4 | × |
| key | 输入 | 公式中的输入 K | 保持与 query、value 的数据类型一致 | FLOAT16、BFLOAT16、INT8 | ND | 3-4 | × |
| value | 输入 | 公式中的输入 V | 保持与 query、key 的数据类型一致 | FLOAT16、BFLOAT16、INT8 | ND | 3-4 | × |
| pseShift | 输入 | 位置编码 | 不使用可传 nullptr;综合约束见约束说明 | FLOAT16、BFLOAT16 | ND | 4 | × |
| attenMask | 输入 | mask 矩阵 | 不使用可传 nullptr;综合约束见约束说明 | BOOL、INT8、UINT8 | ND | 2-4 | × |
| sabi | 输入 | 块稀疏索引矩阵 | 不使用可传 nullptr;语义详见本文 sabi 小节 | UINT16 | ND | 4 | × |
| actualSeqLengths | 输入 | 不同 Batch 中 query 的有效序列长度 | 不指定可传 nullptr;综合约束见约束说明 | INT64 | TND | 1 | - |
| actualSeqLengthsKv | 输入 | 不同 Batch 中 key/value 的有效序列长度 | 不指定可传 nullptr;综合约束见约束说明 | INT64 | TND | 1 | - |
| deqScale1 | 输入 | BMM1 后面的反量化因子 | 支持 per-tensor;不使用可传 nullptr | UINT64、FLOAT32 | ND | 1 | - |
| quantScale1 | 输入 | BMM2 前面的量化因子 | 支持 per-tensor;不使用可传 nullptr | UINT64、FLOAT32 | ND | 1 | - |
| deqScale2 | 输入 | BMM2 后面的反量化因子 | 支持 per-tensor;不使用可传 nullptr | UINT64、FLOAT32 | ND | 1 | - |
| quantScale2 | 输入 | 输出的量化因子 | 支持 per-tensor、per-channel;不使用可传 nullptr | UINT64、FLOAT32 | ND | 1 | - |
| quantOffset2 | 输入 | 输出的量化偏移 | 支持 per-tensor、per-channel;不使用可传 nullptr | FLOAT32 | ND | 1 | - |
| numHeads | 输入 | query 的 head 个数 | - | INT64 | ND | 1 | - |
| scaleValue | 输入 | 公式中 d 开根号的倒数 | 数据类型与 query 满足数据类型推导规则;用户不特意指定时建议传入 1.0 | DOUBLE | - | 1 | - |
| preTokens | 输入 | attention 需要和前几个 Token 计算关联 | 不特意指定时建议传入 2147483647 | INT64 | - | 1 | - |
| nextTokens | 输入 | attention 需要和后几个 Token 计算关联 | 不特意指定时建议传入 0 | INT64 | - | 1 | - |
| inputLayout | 输入 | 标识输入 query、key、value 的数据排布格式 | 不特意指定时建议传入 "BSH";综合约束见约束说明 | CHAR | - | - | - |
| numKeyValueHeads | 输入 | key、value 中 head 个数 | 不特意指定时建议传入 0(表示与 query 相等);综合约束见约束说明 | INT64 | - | - | - |
| sparseMode | 输入 | sparse 的模式 | 综合约束见约束说明 | INT64 | - | 1 | - |
| innerPrecise | 输入 | 高精度或者高性能选择 | 综合约束见约束说明 | INT8 | - | 1 | - |
| attentionOut | 输出 | 公式中的输出 | - | FLOAT16、BFLOAT16、INT8 | ND | 3-4 | - |
| workspaceSize | 输出 | 返回用户需要在 Device 侧申请的 workspace 大小 | - | - | - | 1 | - |
| executor | 输出 | 返回 op 执行器,包含算子计算流程 | - | - | - | 1 | - |
返回值:返回aclnnStatus状态码,具体参见 aclnn返回码。第一段接口完成入参校验,若出现以下错误码,对应原因为:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入参数是必选输入、输出或必选属性且为空指针时返回。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | query、key、value、pseShift、attenMask、attentionOut 的数据类型和数据格式不在支持范围内。 |
| ACLNN_ERR_RUNTIME_ERROR | 361001 | API 内部调用 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 乘积较大)的场景包括但不限于:
| B | Q_N | Q_S | D | KV_N | KV_S |
|---|---|---|---|---|---|
| 1 | 20 | 2097152 | 256 | 1 | 2097152 |
| 1 | 2 | 20971520 | 256 | 2 | 20971520 |
| 20 | 1 | 2097152 | 256 | 1 | 2097152 |
| 1 | 10 | 2097152 | 512 | 1 | 2097152 |
- 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_S、B, Q_S, KV_S、1, Q_S, KV_S、B, 1, Q_S, KV_S、1, 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 > 0或nextTokens < 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)→ExecuteBlitzSparseAttention→aclrtSynchronizeStream→ProcessResults(Device 到 Host 拷贝并打印)→CleanupResources(销毁 tensor、释放内存、销毁 stream、aclFinalize)。
原文档附带的另一版调用示例位于 aclnnBlitzSparseAttention.md 的"调用示例"小节,其使用 BNSD 排布、sparseMode=0、innerPrecise=1、numHeads=2、scaleValue=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.550825Python 侧快速验证: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=0、HIGH_PERFORMANCE=1、MAX_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),仅供参考