接口功能:aclnnBSASelectBlockMask是BSA(BlockSparseAttention)的前置算子,负责根据Query和Key的内容动态生成blockSparseMask,使BSA的调用链从"手动提供掩码"变为"根据Q/K内容自适应选择稀疏模式"。
计算公式:
设blockShape = [blockShapeX, blockShapeY],Sq是query最大序列长度,Skv是key最大序列长度, 则压缩后块数:
Step1:均值池化压缩 (Mean Pooling Compression)
当actualBlockLenQuery / actualBlockLenKey为null时(完整压缩):
当actualBlockLenQuery / actualBlockLenKey非null时(部分压缩),仅对每个block内前actualBlockLen个token取均值:
Step2a:QK Matmul
Step2b:Softmax
Step3:TopK选择生成索引
其中indices为attn_score[b, n, x, y] 中topk_value个最大值对应的索引集合。
Step4:生成BlockSparseMask
数据排布格式:
BSASelectBlockMask输入query、key的数据排布格式支持从多种维度排布解读,可通过qInputLayout和kvInputLayout传入。为了方便理解后续支持的具体排布格式(如BNSD、TND等),此处先对排布格式中各缩写字母所代表的维度含义进行统一说明:
- B:表示输入样本批量大小(Batch)
- T:B和S合轴紧密排列的长度(Total tokens)
- S:表示输入样本序列长度(Seq-Length)
- H:表示隐藏层的大小(Head-Size)
- N:表示多头数(Head-Num)
- D:表示隐藏层最小的单元尺寸,需满足D = H / N(Head-Dim)
当前支持的布局:
- qInputLayout: "TND" "BNSD"
- kvInputLayout: "TND" "BNSD"
每个算子分为,必须先调用"aclnnBSASelectBlockMaskGetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnBSASelectBlockMask"接口执行计算。
- 该接口若与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。
- actualSeqLengths在qInputLayout为 "TND" 时必选;actualSeqLengthsKV在kvInputLayout为 "TND" 时必选。
- 根据算子支持的输入Layout,query张量Shape中对应的head维度大小记为N1,key张量Shape中对应的head维度大小记为N2。必须满足N1 = N2(仅支持MHA)。
- headDim = 128。
- blockShapeX和blockShapeY必须为64的倍数。
- query和key压缩后,query和key对应的Xblocks和Yblocks需满足Xblocks - Yblocks > 1。
- query和key的数据类型必须一致,仅支持FLOAT16和BFLOAT16。
- blockSparseMaskOut数据类型为INT8(二值:0或1)。
- postBlockShape当前不支持,必须传入nullptr。
- actualBlockLenQuery / actualBlockLenKey若非null,每个元素取值范围 [0, blockShapeX] / [0, blockShapeY];为null时完整压缩。
- 不涉及确定性计算。