- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
GenericBlockSparseAttention 是 CANN ops-transformer 中基于 CATLASS 模板库实现的高性能块稀疏注意力算子,支持沿 S 轴(序列维)任意粒度的块稀疏模式,并通过 Paged KV Cache(PA_BBND)与 Packed GQA 特性服务于长序列大模型推理与训练场景。阅读本文后,你将掌握该算子的稀疏分块原理、metadata 前置调度流程、aclnn 两段式 C++ 接口与 PyTorch 扩展接口的完整调用方法,以及全部参数、约束与量化/掩码配置的实操细节。
产品支持情况
GenericBlockSparseAttention 算子及配套的 metadata 前置算子在以下产品上获得支持:
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
产品支持信息同时体现在算子定义中:在 generic_block_sparse_attention_def.cpp 中通过AICore().AddConfig("ascend910b")、AddConfig("ascend910_93")、AddConfig("ascend950")注册了对应芯片的 AICore 配置。
功能说明
稀疏注意力计算机制
GenericBlockSparseAttention 沿序列维(S 轴)按块做稀疏计算。Q 侧按blockShapeX、KV 侧按blockShapeY划分稀疏块,稀疏块大小为:
$$blockShapeX \times blockShapeY$$
sparseBlockIdx:指定每个 Q 块实际选择的 KV 块索引;sparseBlockCount:指定每个 Q 块实际保留的 KV 块数量。
计算时只对选中的块执行 $qk^{T}$、Softmax 以及与 $v$ 的乘积,计算公式为:
$$ attentionOut = Softmax(scaleValue \cdot query \cdot key_{sparse}^{T} + atten_mask) \cdot value_{sparse} $$
其中 $key_{sparse}$、$value_{sparse}$ 为按sparseBlockIdx/sparseBlockCount从 Paged KV Cache 中选取的 KV 块。这种"每个 Q 块只取少量 KV 块"的方式,相比稠密 Attention 大幅减少了参与计算的数据量,是长序列场景下控制计算与访存开销的关键手段。
两段式调用:metadata 前置算子
该算子采用"metadata 生成 + 主算子执行"的两段式架构,这是其与普通 Flash Attention 算子最大的差异点:
- 准备
query、key、value、sparseBlockIdx、sparseBlockCount等输入; - 先调用
aclnnGenericBlockSparseAttentionMetadata(PyTorch 侧为generic_block_sparse_attention_metadata)生成metadataOptional; - 再调用
aclnnGenericBlockSparseAttention(PyTorch 侧为generic_block_sparse_attention),将上一步得到的metadataOptional传入主算子。
metadata 记录的是 AICore/AIVCore 的任务切分结果(即负载均衡调度信息),主算子传入后可以优化任务调度。从 generic_block_sparse_attention.py 的源码可以看到,metadata 是一个 shape 固定为(1024,)的 int32 Tensor(GBSA_METADATA_SIZE = 1024)。
实现架构与 Kernel 侧支撑
从源码结构看,算子内核在 op_kernel 目录下按芯片架构分目录组织:
arch22/:Atlas A2 系列(AICore 220)的内核实现;arch35/:Atlas A3/Ascend 950 系列(AICore 310)的内核实现,包含完整量化路径(generic_block_sparse_attention_kernel_arch35_full_quant.h);attn_infra/:跨架构复用的注意力基础设施,包括 gemm(QK/PV 分块计算)、epilogue(online softmax、rescale、结果写出)、layout 与坐标管理等子模块。
主内核入口 generic_block_sparse_attention.cpp 通过TILING_KEY在编译期展开不同组合(FP16/BF16/FP8、LSE 是否输出、softmaxPrecision 等)。例如 arch22 上GBSA_FP16_TND_PA_BBND_TILING对应softmaxPrecision=0(fp32 Softmax + Rescale),GBSA_FP16_TND_PA_BBND_HALFSM_TILING对应softmaxPrecision=1(half Softmax + fp32 Rescale),这与下文的精度配置一一对应。
参数说明
说明:参数维度含义——B 表示 Batch Size,T 表示 Total tokens,N 表示 Head Num,D 表示 Head Dim,topK 表示
sparseBlockIdx最后一维maxSparseBlockCount。TND 中的 N 为 query 的 headNum(记为 N1),PA_BBND 中的 N 为 key/value 的 headNum(记为 N2),GQA 下 N1 与 N2 可以不同(约束见下文)。
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| query | 输入 | 公式中的 query。layoutQ 为 "TND" 时,shape 为 [T, N, D],N 为 query 的 headNum(N1) | FLOAT16、BFLOAT16、FLOAT8_E4M3FN | ND |
| key | 输入 | 公式中的 key。layoutKv 为 "PA_BBND" 时,shape 为 [numBlocks, blockSize, N, D],N 为 kv 的 headNum(N2) | FLOAT16、BFLOAT16、FLOAT8_E4M3FN | ND |
| value | 输入 | 公式中的 value,shape 与 key 一致 | FLOAT16、BFLOAT16、FLOAT8_E4M3FN | ND |
| sparseBlockIdx | 输入 | 稀疏块索引。TND + isPackedGQA=1 时,shape 为 [N, totalQBlocks, topK],N 为 kv 的 headNum(N2);无效位置可用 -1 填充,有效值须落在前 sparseBlockCount 个位置 | INT32 | ND |
| sparseBlockCount | 输入 | 每个 Q 块实际选择的 KV 块数量。TND + isPackedGQA=1 时,shape 为 [N, totalQBlocks] | INT32 | ND |
| cuSeqLengthsQOptional | 输入 | 各 batch 中 query 序列长度前缀和,layoutQ 为 "TND" 时必传,shape 为 [B+1];第 0 个元素为 0,最后一个元素等于 totalQTokens,相邻差分得到各 batch 的存储长度 | INT64 | ND |
| cuSeqLengthsKvOptional | 输入 | 各 batch 中 key/value 序列长度前缀和,layoutKv 为 "TND" 时必传,非 TND(如 PA_BBND)时不传,shape 为 [B+1] | INT64 | ND |
| sequsedQOptional | 可选输入 | 各 batch 中 query 实际有效长度;不传时按 cu 前缀和差分得到的存储长度处理,shape 为 [B] | INT32 | ND |
| sequsedKvOptional | 输入 | 各 batch 中 kv 实际有效长度,layoutKv 为 "PA_BBND" 时必传,shape 为 [B] | INT32 | ND |
| blockTableOptional | 输入 | PagedAttention 页表,shape 为 [B, maxNumBlocksPerBatch],值只能为正整数 | INT32 | ND |
| blockShape | 属性 | 稀疏块形状 [blockShapeX, blockShapeY],当前仅支持 [1, 128];blockShapeX 支持任意值,blockShapeY 支持按 16 对齐的任意值(均不超过 int64 范围) | INT64 | - |
| isPackedGQA | 属性 | 同 group 内 qHead 是否共享稀疏 pattern,当前仅支持 1(True) | INT64 | - |
| layoutQ | 属性 | query 数据排布格式,目标支持 "TND"/"BNSD"/"BSND",当前仅支持 "TND" | STRING | - |
| layoutKv | 属性 | key/value 数据排布格式,目标支持 "TND"/"BNSD"/"BSND"/"PA_BBND"/"PA_BNBD",当前仅支持 "PA_BBND" | STRING | - |
| layoutSparsePattern | 属性 | sparseBlockIdx、sparseBlockCount 的数据排布格式,当前仅支持取 4 | INT64 | - |
| scaleValue | 属性 | 缩放系数;传 0 时算子内按 $1/\sqrt{D}$ 处理,一般设置为 D^-0.5 | DOUBLE | - |
| maskType | 属性 | 掩码类型,取值 0~5,当前仅支持 1(内置 causal mask) | INT64 | - |
| quantType | 属性 | 量化类型;当前支持 0,Ascend 950 上可选 5,取值 1~4 传入将校验失败 | INT64 | - |
| dstTypeMax | 属性 | MXFP4 CX 量化时传入的自定义量化量程,当前版本不支持自定义量程,必须传入 0.0 | DOUBLE | - |
| softmaxPrecision | 属性 | Softmax 计算精度级别,取值 0 或 1(详见下文"Softmax 精度") | INT64 | - |
| winLeft / winRight | 属性 | 滑窗 attention 场景的前向/后向窗口 token 数;当前不支持滑窗,只支持传入 -1 | INT64 | - |
| residualBlockMode | 属性 | KV 序列按 blockShapeY 稀疏后尾部不完整块的状态,仅支持 0 或 1(见下文约束) | INT64 | - |
| isConsistentTopK | 属性 | 同一 batch 同一 head 内每个 Q 块选择的 KV 块最大数量是否一致,仅支持 0 或 1 | BOOL | - |
| returnSoftmaxlse | 属性 | 是否输出 softmaxLse,当前仅支持 0 | INT64 | - |
| attentionOut | 输出 | 公式中的 attentionOut,数据类型和 shape 与 query 保持一致;FP8 输入时由本 tensor 指定输出 dtype | FLOAT16、BFLOAT16 | ND |
| softmaxLseOptional | 输出 | Softmax log-sum-exp 中间结果;当前不支持,须传入 nullptr(returnSoftmaxlse 须为 0) | FLOAT | ND |
在算子定义文件 generic_block_sparse_attention_def.cpp 中可以看到这些属性的默认值:block_shape默认[1, 128]、is_packed_gqa默认 1、layout_q默认 "TND"、scale_value默认 0.0、mask_type默认 0、quant_type默认 0、win_left/win_right默认 -1,与上文表格一致。同时该文件还揭示了非连续 Tensor 的处理策略:key/value使用IgnoreContiguous()(允许非连续输入),而query、sparseBlockIdx等使用AutoContiguous()(自动做连续化)。
约束说明
通用约束
- metadata 前置要求:调用前须先执行
aclnnGenericBlockSparseAttentionMetadata生成metadataOptional,再调用本接口;metadata 须与当前输入/属性配套,每次调用须重新生成。metadata 接口与主算子的 sparseBlockIdx、sparseBlockCount、blockShape、mask_mode、quant_mode、is_packed_gqa 等参数须完全一致,否则产生未定义行为(精度问题或非法内存访问导致的崩溃)。 - 维度约束:query/key/value 的 headDim(D)当前仅支持 128;KV 页 blockSize 当前仅支持 128,且须等于 blockShapeY。
- 分块与寻址:TND + isPackedGQA=1 时,totalQBlocks 按 cuSeqLengthsQ 差分得到的存储长度分块($\sum_i \mathrm{ceilDiv}(qStorageLen_i, blockShapeX)$);sparse 分块与 QKV 寻址均按该存储长度,不以 seqused 重切分;topK 须 ≥ sparseBlockCount 中所有元素的最大值,当前上限为 256。
- 双长度语义:seqused 与对应 cu 前缀和同时传入时(Q 侧,及 layoutKv 为 TND 时的 KV 侧),分核/任务空间按各 batch 实际有效长度(seqused)累加;各 batch 的 seqused 元素须 ≤ 对应 cu 存储长度,且须与 Metadata 侧完全一致。
- PA_BBND 特有约束:layoutKv 为 "PA_BBND" 时须传
sequsedKvOptional,不传cuSeqLengthsKvOptional。 - 数据类型一致性:输入 query、key、value 的数据类型必须一致。
- GQA 约束:query 的 headNum 为 N1,key/value 的 headNum 为 N2,则 N1 ≥ N2 且 N1 % N2 == 0;groupSize = N1/N2 当前须 ≤ 128。
- 非连续约束:PA_BBND 下 key/value 仅 dim0(物理页轴)可非连续;页内 blockSize × N2 × D 须连续,且 stride0 ≥ blockSize × N2 × D 并按 N2 × D 对齐。
- 确定性:aclnnGenericBlockSparseAttention 默认确定性实现。
- 空 Tensor:必选输入和输出 shape 中任意轴为 0 的空 Tensor 用例将全部拦截报错。
- Tiling 期不校验值:cu_seqlens、seqused、sparseBlockIdx、sparseBlockCount 及 blockTable 等 Tensor 在 Tiling 阶段无法获取具体数值,tiling 侧不对其值进行校验,正确性需要用户自行保证。
Softmax 精度(softmaxPrecision)
softmaxPrecision 控制 online softmax 阶段以及 rescale 阶段运算使用的数据类型:
- 0:online softmax 和 rescale 全部采取 fp32,适合追求计算精度的场景;
- 1:混合精度;online softmax 采取 fp16/bf16(与 attentionOut 相同),rescale 采取 fp32,online softmax 阶段可能数值溢出。
芯片约束:Ascend 950 仅支持 1;Atlas A2/A3 上 FLOAT16 可配置 0 或 1,BFLOAT16 仅支持 0;FP8 路径仅支持 1。内核侧的对应关系可参见 generic_block_sparse_attention.cpp 中TILING_KEY的注释。
掩码说明(maskType)
| maskType | 含义 | attentionMaskOptional | winLeft/winRight |
|---|---|---|---|
| 0 | 不加 mask | 不传 | -1/-1 |
| 1 | causal mask | 不传 attenMaskOptional(内置 causal) | -1/-1 |
| 2 | window mask | 规划:BOOL [2048,2048] 下三角,与 winLeft/winRight 配合 | 实际 window 包含的向前/向后看 token 数 |
| 3~5 | 各类特化 mask | 后续补充 mask 描述 | -1/-1 |
当前仅支持 maskType=1(算子内置 causal,attenMaskOptional 须为 nullptr,winLeft/winRight 为 -1);maskType 为 0/2/3~5 当前不支持。
量化说明(quantType)
量化配置按 Ascend 950(A5)代际描述,完整配置如下:
| quantType | QKV 数据类型 | 对称/非对称 | P 量化动态/静态 | 量化粒度 | 量化参数 shape | 量化参数 dType |
|---|---|---|---|---|---|---|
| 0 | 非量化,QKV 直接作为输入计算 | - | - | - | q/k/vDequantScaleOptional、pQuantScaleOptional 均不传 | - |
| 1 | FLOAT8_E4M3 | 对称 | 静态 | perGroup,QKV 均沿 S 维度分组,group 大小和稀疏块尺寸必须相同;KV 为 paged cache 时 blockSize 需为 blockShapeY 的整数倍 | q/k/vDequantScaleOptional 必选,pQuantScaleOptional 可选(传入为 [1] 静态系数,nullptr 时默认 448.0) | FLOAT32 |
| 2 | FLOAT8_E4M3 | 对称 | 动态 | micro scaling,QKV 沿矩阵乘累加轴按固定大小 32 分组;KV 为 paged cache 时 blockSize 需为 64 的整数倍 | q/k/vDequantScaleOptional 必选 | FLOAT8_E4M3 |
| 3 | FLOAT4_E2M1 | 对称 | 动态 OCP | 同 quantType=2 | 同 quantType=2 | FLOAT8_E4M3 |
| 4 | FLOAT4_E2M1 | 对称 | 动态 CX | 同 quantType=2 | 同 quantType=2 | FLOAT8_E4M3 |
| 5 | FLOAT8_E4M3 | 对称 | 静态 | 不传入量化系数,算子内直接将 P cast 成 fp8 | q/k/vDequantScaleOptional、pQuantScaleOptional 均不传 | FLOAT32 |
当前可用 quantType=0;Ascend 950 上可选 quantType=5。quantType=1~4 当前不支持,传入将校验失败。quantType=0 时 q/k/v_dequant_scale 须为 None;quantType≠0 时 attention_out_dtype 必须传入(量化场景输出 dtype 须为 float16 或 bfloat16)。
在 arch35 内核中,量化路径由GbsaInferInterfaceFullQuant<fp8_e4m3fn_t, ...>模板实例化,对应 generic_block_sparse_attention_kernel_arch35_full_quant.h,实现了 FP8 输入的完整反量化与计算流程。
调用说明
GenericBlockSparseAttention 提供两套调用方式:
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn API | test_aclnn_generic_block_sparse_attention.cpp | 通过 aclnnGenericBlockSparseAttention 两段式接口调用 GenericBlockSparseAttention 算子 |
| PyTorch API | 见下文 Python 示例 | 通过 generic_block_sparse_attention 接口调用 generic_block_sparse_attention 算子 |
aclnn API:两段式接口调用
aclnn 接口为两段式设计:先调用aclnnGenericBlockSparseAttentionGetWorkspaceSize获取 workspace 大小与执行器,再调用aclnnGenericBlockSparseAttention执行计算。
第一段接口原型:
aclnnStatus aclnnGenericBlockSparseAttentionGetWorkspaceSize( const aclTensor *query, const aclTensor *key, const aclTensor *value, const aclTensor *sparseBlockIdx, const aclTensor *sparseBlockCount, const aclTensor *metadataOptional, const aclTensor *attenMaskOptional, const aclTensor *qDequantScaleOptional, const aclTensor *kDequantScaleOptional, const aclTensor *vDequantScaleOptional, const aclTensor *pQuantScaleOptional, const aclTensor *cuSeqLengthsQOptional, const aclTensor *cuSeqLengthsKvOptional, const aclTensor *sequsedQOptional, const aclTensor *sequsedKvOptional, const aclTensor *blockTableOptional, const aclIntArray *blockShape, char *layoutQ, char *layoutKv, int64_t layoutSparsePattern, double scaleValue, int64_t maskType, int64_t quantType, double dstTypeMax, int64_t softmaxPrecision, int64_t winLeft, int64_t winRight, int64_t returnSoftmaxlse, int64_t residualBlockMode, bool isConsistentTopK, aclTensor *attentionOut, aclTensor *softmaxLseOptional, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口原型:
aclnnStatus aclnnGenericBlockSparseAttention( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)典型的 C++ 调用流程(完整示例见 examples/test_aclnn_generic_block_sparse_attention.cpp 与 docs/aclnnGenericBlockSparseAttention.md):
#include "acl/acl.h" #include "aclnnop/aclnn_generic_block_sparse_attention.h" #include "aclnnop/aclnn_generic_block_sparse_attention_metadata.h" // 1. 构造输入/输出 aclTensor(query: [T,N1,D] fp16;key/value: [numBlocks,blockSize,N2,D] fp16; // sparseBlockIdx: [N2,totalQBlocks,topK] int32;sparseBlockCount: [N2,totalQBlocks] int32; // cuSeqLengthsQ: [B+1] int64;sequsedKv: [B] int32;blockTable: [B,maxBlocks] int32) // 构造方式见样例中的 CreateAclTensor 辅助函数(aclrtMalloc + aclrtMemcpy + aclCreateTensor) // 2. 先调用 Metadata 算子生成任务切分结果 uint64_t metadataWorkspaceSize = 0; aclOpExecutor* metadataExecutor = nullptr; ret = aclnnGenericBlockSparseAttentionMetadataGetWorkspaceSize( sparseIdx, sparseCount, cuSeqQ, nullptr, nullptr, sequsedKv, S1, S2, N1, N2, D, blockShape, layoutQ, layoutKv, /*layoutSparsePattern=*/4, /*maskType=*/1, /*quantType=*/0, /*softmaxPrecision=*/1, -1, -1, 0, 0, metadata, &metadataWorkspaceSize, &metadataExecutor); // 申请 workspace 后执行: ret = aclnnGenericBlockSparseAttentionMetadata(metadataWorkspaceAddr, metadataWorkspaceSize, metadataExecutor, stream); ret = aclrtSynchronizeStream(stream); // 3. 调用主算子 GetWorkspaceSize 并执行 uint64_t workspaceSize = 0; aclOpExecutor* executor = nullptr; ret = aclnnGenericBlockSparseAttentionGetWorkspaceSize( q, k, v, sparseIdx, sparseCount, metadata, /*attenMask=*/nullptr, /*qDequantScale=*/nullptr, /*kDequantScale=*/nullptr, /*vDequantScale=*/nullptr, /*pQuantScale=*/nullptr, cuSeqQ, /*cuSeqLengthsKv=*/nullptr, /*sequsedQ=*/nullptr, sequsedKv, blockTable, blockShape, layoutQ, layoutKv, /*layoutSparsePattern=*/4, scaleValue, /*maskType=*/1, /*quantType=*/0, /*dstTypeMax=*/0.0, /*softmaxPrecision=*/1, /*winLeft=*/-1, /*winRight=*/-1, /*returnSoftmaxlse=*/0, /*residualBlockMode=*/0, /*isConsistentTopK=*/0, attnOut, /*softmaxLse=*/nullptr, &workspaceSize, &executor); // 申请 workspace 后执行: ret = aclnnGenericBlockSparseAttention(workspaceAddr, workspaceSize, executor, stream); ret = aclrtSynchronizeStream(stream);第一段接口完成入参校验,常见返回码如下:
| 返回码 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | query/key/value/sparseBlockIdx/sparseBlockCount/attentionOut 等必选指针为空 |
| ACLNN_ERR_PARAM_INVALID | 161002 | layout、maskType、blockShape、softmaxPrecision、quantType、returnSoftmaxlse、layoutSparsePattern、residualBlockMode、isConsistentTopK 等与约束不匹配 |
| ACLNN_ERR_INNER_NULLPTR | 561103 | metadata 为空或 Contiguous/InferShape 失败(如 layout 不支持、缺少 blockTable 等) |
PyTorch API:TorchNPU 扩展接口
PyTorch 侧接口定义于 torch_extension/generic_block_sparse_attention.py,通过 TorchNPU 的cann_ops_transformer扩展以torch.library机制注册(torch.ops.cann_ops_transformer.generic_block_sparse_attention)。
前置 metadata 接口:
cann_ops_transformer.generic_block_sparse_attention_metadata( sparse_block_idx, sparse_block_count, num_heads_q, num_heads_kv, head_dim, block_shape, *, cu_seqlens_q=None, cu_seqlens_kv=None, seqused_q=None, seqused_kv=None, max_seqlen_q=-1, max_seqlen_kv=-1, is_packed_gqa=True, layout_q="TND", layout_kv="PA_BBND", mask_mode=1, quant_mode=0, softmax_precision=1, win_left=-1, win_right=-1, ) -> Tensor主算子接口:
cann_ops_transformer.generic_block_sparse_attention( q, k, v, sparse_block_idx, sparse_block_count, block_shape, *, metadata=None, attn_mask=None, q_dequant_scale=None, k_dequant_scale=None, v_dequant_scale=None, p_quant_scale=None, cu_seqlens_q=None, cu_seqlens_kv=None, seqused_q=None, seqused_kv=None, block_table=None, is_packed_gqa=True, layout_q="TND", layout_kv="PA_BBND", softmax_scale=0.0, mask_mode=1, quant_mode=0, dst_type_max=0.0, softmax_precision=1, win_left=-1, win_right=-1, return_softmax_lse=False, attention_out_dtype=None, ) -> (Tensor, Tensor)枚举说明
quant_mode与mask_mode在 Python 接口中支持传入IntEnum枚举或对应 int 值,枚举定义于cann_ops_transformer.ops.generic_block_sparse_attention(源码见 generic_block_sparse_attention.py):
quant_mode 枚举(QuantMode)
| 枚举名 | 值 | 含义 |
|---|---|---|
NO_QUANT | 0 | 非量化(默认值) |
FP8_E4M3_STATIC_PER_GROUP | 1 | FP8_E4M3 静态 per-group |
FP8_E4M3_DYNAMIC_MX | 2 | FP8_E4M3 动态 MX |
FP4_E2M1_DYNAMIC_OCP | 3 | FP4_E2M1 动态 OCP |
FP4_E2M1_DYNAMIC_CX | 4 | FP4_E2M1 动态 CX |
FP8_E4M3_STATIC_CAST_P | 5 | FP8_E4M3 静态 cast P |
mask_mode 枚举(MaskMode)
| 枚举名 | 值 | 含义 |
|---|---|---|
NO_MASK | 0 | 不加 mask |
CAUSAL | 1 | Causal 模式(默认值) |
WINDOW | 2 | Window 模式 |
当前仅支持 mask_mode = 1(CAUSAL);quant_mode 当前仅支持 0(NO_QUANT)与 5(FP8_E4M3_STATIC_CAST_P)。
Python 联合调用示例(TND + PA_BBND)
以下示例取自 torchapi_generic_block_sparse_attention.md:
import math import torch import torch_npu import cann_ops_transformer torch_npu.npu.set_device(0) B, Q_N, KV_N, Q_S, KV_S, D = 1, 32, 8, 128, 256, 128 block_x, block_y = 1, 128 block_size = 128 top_k = 16 Q_T = B * Q_S num_blocks = math.ceil(KV_S / block_size) total_q_blocks = math.ceil(Q_S / block_x) q = torch.randn(Q_T, Q_N, D, dtype=torch.float16, device="npu") k = torch.randn(num_blocks, block_size, KV_N, D, dtype=torch.float16, device="npu") v = torch.randn(num_blocks, block_size, KV_N, D, dtype=torch.float16, device="npu") sparse_block_idx = torch.randint( 0, num_blocks, (KV_N, total_q_blocks, top_k), dtype=torch.int32, device="npu" ) sparse_block_count = torch.full((KV_N, total_q_blocks), top_k, dtype=torch.int32, device="npu") cu_seqlens_q = torch.tensor([0, Q_S], dtype=torch.int64, device="npu") seqused_kv = torch.tensor([KV_S], dtype=torch.int32, device="npu") block_table = torch.arange(num_blocks, dtype=torch.int32, device="npu").view(1, -1) block_shape = [block_x, block_y] metadata = cann_ops_transformer.ops.generic_block_sparse_attention_metadata( sparse_block_idx, sparse_block_count, Q_N, KV_N, D, block_shape, cu_seqlens_q=cu_seqlens_q, seqused_kv=seqused_kv, max_seqlen_q=Q_S, max_seqlen_kv=KV_S, is_packed_gqa=True, layout_q="TND", layout_kv="PA_BBND", mask_mode=1, quant_mode=0, softmax_precision=1, ) attention_out, softmax_lse = cann_ops_transformer.ops.generic_block_sparse_attention( q, k, v, sparse_block_idx, sparse_block_count, block_shape, metadata=metadata, cu_seqlens_q=cu_seqlens_q, seqused_kv=seqused_kv, block_table=block_table, is_packed_gqa=True, layout_q="TND", layout_kv="PA_BBND", softmax_scale=1.0 / (D ** 0.5), mask_mode=1, quant_mode=0, softmax_precision=1, return_softmax_lse=False, ) torch_npu.npu.synchronize() assert attention_out.shape == q.shape assert attention_out.dtype == q.dtype返回值说明
generic_block_sparse_attention_metadata返回 shape 为(1024,)的 int32 Tensor(任务切分数据);generic_block_sparse_attention返回attention_out(shape 与 q 一致;quant_mode=0 时 dtype 默认与 q 一致,quant_mode≠0 时由 attention_out_dtype 指定)和softmax_lse(return_softmax_lse=True 时 TND 布局下输出 shape 为(Q_T, Q_N, 1)的 float32 Tensor,否则为空 Tensor)。该逻辑在 generic_block_sparse_attention.py 的register_meta中有对应实现。
Layout 与基准维度速查
基准符号
| 命名 | 含义 |
|---|---|
| B | Batch Size |
| T / totalQTokens | query 的 Total tokens(所有 batch 序列长度累加和) |
| N / headNum | query 的 head 数(N1) |
| numKeyValueHeads | key/value 的 head 数(N2) |
| D / headDim | Head Dim,且满足 D = H / N |
| numBlocks | Paged KV Cache 的物理页数 |
| blockSize | 每一页容纳的 token 数 |
| maxNumBlocksPerBatch | blockTable 第二维,须 ≥ ceilDiv(maxKvSeqLength, blockSize) |
| totalQBlocks | 按存储长度分块后的 Q 块总数,$\sum_i \mathrm{ceilDiv}(qStorageLen_i, blockShapeX)$ |
| totalKBlocks | 按存储长度分块后的 KV 块总数,$\sum_i \mathrm{ceilDiv}(kvStorageLen_i, blockShapeY)$ |
| maxKvBlockCount / topK | sparseBlockIdx 最后一维,须不小于 sparseBlockCount 中所有元素的最大值,当前上限为 256 |
| blockShapeX / blockShapeY | 稀疏块在 Q 方向、KV 方向的块大小 |
layoutSparsePattern 与 sparse 张量 shape
layoutSparsePattern决定 sparseBlockIdx、sparseBlockCount 的 shape(当前仅支持值 4):
| layoutSparsePattern | sparseBlockIdx | sparseBlockCount | 描述 |
|---|---|---|---|
| 0 | [batch, N2, maxQBlockCount, maxKvBlockCount] | [batch, N2, maxQBlockCount] | 同 group qHead 共享 pattern,表示每个 Q 块选了哪些 KV 块 |
| 1 | [batch, N2, maxKvBlockCount, maxQBlockCount] | [batch, N2, maxKvBlockCount] | 同 group qHead 共享 pattern,表示每个 KV 块选了哪些 Q 块 |
| 2 | [batch, N1, maxQBlockCount, maxKvBlockCount] | [batch, N1, maxQBlockCount] | 同 group qHead 独立 pattern,表示每个 Q 块选了哪些 KV 块 |
| 3 | [batch, N1, maxKvBlockCount, maxQBlockCount] | [batch, N1, maxKvBlockCount] | 同 group qHead 独立 pattern,表示每个 KV 块选了哪些 Q 块 |
| 4 | [N2, totalQBlocks, maxKvBlockCount] | [N2, totalQBlocks] | 同 group qHead 共享 pattern,表示每个 Q 块选了哪些 KV 块(当前支持) |
| 5 | [N2, totalKBlocks, maxQBlockCount] | [N2, totalKBlocks] | 同 group qHead 共享 pattern,表示每个 KV 块选了哪些 Q 块 |
| 6 | [N1, totalQBlocks, maxKvBlockCount] | [N1, totalQBlocks] | 同 group qHead 独立 pattern,表示每个 Q 块选了哪些 KV 块 |
| 7 | [N1, totalKBlocks, maxQBlockCount] | [N1, totalKBlocks] | 同 group qHead 独立 pattern,表示每个 KV 块选了哪些 Q 块 |
Paged Attention 相关
| blockTable | kvLayout | Key/Value shape |
|---|---|---|
| 非空,shape 为 [batch, maxNumBlocksPerBatch],代表使能 paged cache | PA_BBND | [numBlocks, blockSize, numKeyValueHeads, headDim] |
| PA_BNBD | [numBlocks, numKeyValueHeads, blockSize, headDim] | |
| 空,代表不使能 paged cache,接收原始 KV | TND / BSND / BNSD | [totalKTokens, N2, D] / [batch, maxKvSeqLength, N2, D] / [batch, N2, maxKvSeqLength, D] |
当前必须传入非空 blockTable,且 layoutKv 为 PA_BBND;blockTable 为 nullptr(原始 KV)及 PA_BNBD 当前不支持。PagedAttention 开启情况下还必须传入 sequsedKv。
源码与测试佐证
- 算子定义:generic_block_sparse_attention_def.cpp 定义了全部输入、输出与属性的注册信息及默认值。
- InferShape:generic_block_sparse_attention_infershape.cpp 负责输出 shape/dtype 推导。
- Tiling:generic_block_sparse_attention_tiling.cpp 实现基于 metadata 的分核与任务切分;单测见 test_generic_block_sparse_attention_tiling.cpp。
- 内核入口:generic_block_sparse_attention.cpp 通过 TILING_KEY 分发到 arch22/arch35 各实现。
- PyTorch 绑定:generic_block_sparse_attention.py 与 csrc/generic_block_sparse_attention.cpp 完成 torch.library 注册与底层调用。
- UT 测试:test_aclnn_generic_block_sparse_attention.cpp 覆盖 aclnn 接口的 op_api 级验证;test_generic_block_sparse_attention_infershape.cpp 覆盖 shape 推导。
总结
GenericBlockSparseAttention 是 CANN ops-transformer 中面向长序列、稀疏注意力场景的高性能算子:它通过sparseBlockIdx/sparseBlockCount描述每个 Q 块对 KV 块的稀疏选择,借助独立的 metadata 前置算子完成负载均衡任务切分,并基于 CATLASS 模板库在 arch22/arch35 两代芯片上提供 FP16/BF16/FP8 计算路径。当前版本聚焦于 TND + PA_BBND 组合(Packed GQA + Paged KV Cache),支持内置 causal mask 与 quantType 0/5 两种量化模式。开发者既可以通过 aclnn 两段式 C++ 接口在异构流程中精细控制 workspace 与执行器,也可以通过 TorchNPU 的cann_ops_transformer.generic_block_sparse_attention接口在 PyTorch 图/命令式流程中直接使用,两份调用示例均可直接在支持产品上运行验证。
- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
相关推荐
CANN ops-transformer 算子详解:aclnnBlockSparseAttentionV2 块稀疏注意力接口
CANN ops transformer 算子详解:aclnnBlockSparseAttentionV2 块稀疏注意力接口 导读 aclnnBlockSpar
算子库人工智能深度学习AscendCANN ops-transformer 中的 RainFusionAttention 算子:块级稀疏注意力原理、ACLNN 两段式接口与实战调用
CANN ops transformer 中的 RainFusionAttention 算子:块级稀疏注意力原理、ACLNN 两段式接口与实战调用 本篇技术指南
算子库人工智能深度学习AscendopenEuler-agreements社区协作指南:如何在AtomGit平台高效贡献
openEuler agreements社区协作指南:如何在AtomGit平台高效贡献 前往项目官网免费下载: https://ar.openeuler.org
算子库人工智能深度学习Ascend
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考