CANN ops-transformer GenericBlockSparseAttention 算子深度指南:基于 CATLASS 的块稀疏注意力实现与 aclnn/PyTorch 双接口调用
2026/9/20 22:38:53 网站建设 项目流程
  • 算子库
  • 人工智能
  • 深度学习
  • Ascend

【免费下载链接】ops-transformer

本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

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 算子最大的差异点:

  1. 准备querykeyvaluesparseBlockIdxsparseBlockCount等输入;
  2. 先调用aclnnGenericBlockSparseAttentionMetadata(PyTorch 侧为generic_block_sparse_attention_metadata)生成metadataOptional
  3. 再调用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_E4M3FNND
key输入公式中的 key。layoutKv 为 "PA_BBND" 时,shape 为 [numBlocks, blockSize, N, D],N 为 kv 的 headNum(N2)FLOAT16、BFLOAT16、FLOAT8_E4M3FNND
value输入公式中的 value,shape 与 key 一致FLOAT16、BFLOAT16、FLOAT8_E4M3FNND
sparseBlockIdx输入稀疏块索引。TND + isPackedGQA=1 时,shape 为 [N, totalQBlocks, topK],N 为 kv 的 headNum(N2);无效位置可用 -1 填充,有效值须落在前 sparseBlockCount 个位置INT32ND
sparseBlockCount输入每个 Q 块实际选择的 KV 块数量。TND + isPackedGQA=1 时,shape 为 [N, totalQBlocks]INT32ND
cuSeqLengthsQOptional输入各 batch 中 query 序列长度前缀和,layoutQ 为 "TND" 时必传,shape 为 [B+1];第 0 个元素为 0,最后一个元素等于 totalQTokens,相邻差分得到各 batch 的存储长度INT64ND
cuSeqLengthsKvOptional输入各 batch 中 key/value 序列长度前缀和,layoutKv 为 "TND" 时必传,非 TND(如 PA_BBND)时不传,shape 为 [B+1]INT64ND
sequsedQOptional可选输入各 batch 中 query 实际有效长度;不传时按 cu 前缀和差分得到的存储长度处理,shape 为 [B]INT32ND
sequsedKvOptional输入各 batch 中 kv 实际有效长度,layoutKv 为 "PA_BBND" 时必传,shape 为 [B]INT32ND
blockTableOptional输入PagedAttention 页表,shape 为 [B, maxNumBlocksPerBatch],值只能为正整数INT32ND
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 的数据排布格式,当前仅支持取 4INT64-
scaleValue属性缩放系数;传 0 时算子内按 $1/\sqrt{D}$ 处理,一般设置为 D^-0.5DOUBLE-
maskType属性掩码类型,取值 0~5,当前仅支持 1(内置 causal mask)INT64-
quantType属性量化类型;当前支持 0,Ascend 950 上可选 5,取值 1~4 传入将校验失败INT64-
dstTypeMax属性MXFP4 CX 量化时传入的自定义量化量程,当前版本不支持自定义量程,必须传入 0.0DOUBLE-
softmaxPrecision属性Softmax 计算精度级别,取值 0 或 1(详见下文"Softmax 精度")INT64-
winLeft / winRight属性滑窗 attention 场景的前向/后向窗口 token 数;当前不支持滑窗,只支持传入 -1INT64-
residualBlockMode属性KV 序列按 blockShapeY 稀疏后尾部不完整块的状态,仅支持 0 或 1(见下文约束)INT64-
isConsistentTopK属性同一 batch 同一 head 内每个 Q 块选择的 KV 块最大数量是否一致,仅支持 0 或 1BOOL-
returnSoftmaxlse属性是否输出 softmaxLse,当前仅支持 0INT64-
attentionOut输出公式中的 attentionOut,数据类型和 shape 与 query 保持一致;FP8 输入时由本 tensor 指定输出 dtypeFLOAT16、BFLOAT16ND
softmaxLseOptional输出Softmax log-sum-exp 中间结果;当前不支持,须传入 nullptr(returnSoftmaxlse 须为 0)FLOATND

在算子定义文件 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()(允许非连续输入),而querysparseBlockIdx等使用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含义attentionMaskOptionalwinLeft/winRight
0不加 mask不传-1/-1
1causal mask不传 attenMaskOptional(内置 causal)-1/-1
2window 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)代际描述,完整配置如下:

quantTypeQKV 数据类型对称/非对称P 量化动态/静态量化粒度量化参数 shape量化参数 dType
0非量化,QKV 直接作为输入计算---q/k/vDequantScaleOptional、pQuantScaleOptional 均不传-
1FLOAT8_E4M3对称静态perGroup,QKV 均沿 S 维度分组,group 大小和稀疏块尺寸必须相同;KV 为 paged cache 时 blockSize 需为 blockShapeY 的整数倍q/k/vDequantScaleOptional 必选,pQuantScaleOptional 可选(传入为 [1] 静态系数,nullptr 时默认 448.0)FLOAT32
2FLOAT8_E4M3对称动态micro scaling,QKV 沿矩阵乘累加轴按固定大小 32 分组;KV 为 paged cache 时 blockSize 需为 64 的整数倍q/k/vDequantScaleOptional 必选FLOAT8_E4M3
3FLOAT4_E2M1对称动态 OCP同 quantType=2同 quantType=2FLOAT8_E4M3
4FLOAT4_E2M1对称动态 CX同 quantType=2同 quantType=2FLOAT8_E4M3
5FLOAT8_E4M3对称静态不传入量化系数,算子内直接将 P cast 成 fp8q/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 APItest_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_NULLPTR161001query/key/value/sparseBlockIdx/sparseBlockCount/attentionOut 等必选指针为空
ACLNN_ERR_PARAM_INVALID161002layout、maskType、blockShape、softmaxPrecision、quantType、returnSoftmaxlse、layoutSparsePattern、residualBlockMode、isConsistentTopK 等与约束不匹配
ACLNN_ERR_INNER_NULLPTR561103metadata 为空或 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_modemask_mode在 Python 接口中支持传入IntEnum枚举或对应 int 值,枚举定义于cann_ops_transformer.ops.generic_block_sparse_attention(源码见 generic_block_sparse_attention.py):

quant_mode 枚举(QuantMode)

枚举名含义
NO_QUANT0非量化(默认值)
FP8_E4M3_STATIC_PER_GROUP1FP8_E4M3 静态 per-group
FP8_E4M3_DYNAMIC_MX2FP8_E4M3 动态 MX
FP4_E2M1_DYNAMIC_OCP3FP4_E2M1 动态 OCP
FP4_E2M1_DYNAMIC_CX4FP4_E2M1 动态 CX
FP8_E4M3_STATIC_CAST_P5FP8_E4M3 静态 cast P

mask_mode 枚举(MaskMode)

枚举名含义
NO_MASK0不加 mask
CAUSAL1Causal 模式(默认值)
WINDOW2Window 模式

当前仅支持 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 与基准维度速查

基准符号

命名含义
BBatch Size
T / totalQTokensquery 的 Total tokens(所有 batch 序列长度累加和)
N / headNumquery 的 head 数(N1)
numKeyValueHeadskey/value 的 head 数(N2)
D / headDimHead Dim,且满足 D = H / N
numBlocksPaged KV Cache 的物理页数
blockSize每一页容纳的 token 数
maxNumBlocksPerBatchblockTable 第二维,须 ≥ ceilDiv(maxKvSeqLength, blockSize)
totalQBlocks按存储长度分块后的 Q 块总数,$\sum_i \mathrm{ceilDiv}(qStorageLen_i, blockShapeX)$
totalKBlocks按存储长度分块后的 KV 块总数,$\sum_i \mathrm{ceilDiv}(kvStorageLen_i, blockShapeY)$
maxKvBlockCount / topKsparseBlockIdx 最后一维,须不小于 sparseBlockCount 中所有元素的最大值,当前上限为 256
blockShapeX / blockShapeY稀疏块在 Q 方向、KV 方向的块大小

layoutSparsePattern 与 sparse 张量 shape

layoutSparsePattern决定 sparseBlockIdx、sparseBlockCount 的 shape(当前仅支持值 4):

layoutSparsePatternsparseBlockIdxsparseBlockCount描述
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 相关

blockTablekvLayoutKey/Value shape
非空,shape 为 [batch, maxNumBlocksPerBatch],代表使能 paged cachePA_BBND[numBlocks, blockSize, numKeyValueHeads, headDim]
PA_BNBD[numBlocks, numKeyValueHeads, blockSize, headDim]
空,代表不使能 paged cache,接收原始 KVTND / 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上加速计算。

项目地址:https://gitcode.com/cann/ops-transformer
点击查看免费下载

相关推荐

上一篇:彻底解放Mac生产力:AeroSpace中exec-and-forget命令的避坑指南
下一篇:2025超全Actual Budget开发环境搭建指南:从源码到桌面应用全流程

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

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

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

立即咨询