接口功能:BlockSparseAttentionV2稀疏注意力计算,支持灵活的块级稀疏模式,通过BlockSparseMask指定每个Q块选择的KV块,实现高效的稀疏注意力计算。
相比于BlockSparseAttention,本接口新增qDequantScaleOptional、kDequantScaleOptional、vDequantScaleOptional参数。
计算公式:稀疏块大小:,selectIdx指定稀疏模式
BlockSparseAttentionV2输入query、key、value的数据排布格式支持从多种维度排布解读,可通过qInputLayout和kvInputLayout传入。
- 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" "BSND"
- kvInputLayout: "TND" "BNSD" "BSND"
FP8特性说明(仅[object Object]Ascend 950PR/Ascend 950DT[object Object]支持):
本算子新增支持FP8数据类型的输入,以提供计算效率并降低显存占用。当使用FP8输入时,需要提供相应的量化缩放因子用于反量化计算。
量化缩放因子:
当输入的query、key、value采用FLOAT8_E4M3FN数据类型时,需要提供以下量化缩放因子参数:
qDequantScale(query量化缩放因子)
- 数据类型:FLOAT32
- shape:(Batch, HeadNum, CeilDiv(maxQSeqLength, 128), 1)
- 用途:在QK矩阵乘法时对query进行反量化。
kDequantScale(key量化缩放因子)
- 数据类型:FLOAT32
- shape:(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 256), 1)或(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 512), 1)
- 用途:在QK矩阵乘法时对key进行反量化。
vDequantScale(value量化缩放因子)
- 数据类型:FLOAT32
- shape:(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 256), 1)或(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 512), 1)
- 用途:在PV矩阵乘法时对value进行反量化。
每个算子分为,必须先调用"aclnnBlockSparseAttentionV2GetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnBlockSparseAttentionV2"接口执行计算。
[object Object]
[object Object]
- 确定性计算:
- aclnnBlockSparseAttentionV2默认确定性实现。
- 该接口与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。
- qInputLayout当前仅支持"TND"和"BNSD"和"BSND"。
- kvInputLayout当前仅支持"TND"和"BNSD"和"BSND"。
- 当前query、key、value的InputLayout必须保持一致。
- 输入query、key、value的数据类型必须一致,支持FLOAT16和BFLOAT16。
- query、key、value的D轴当前仅支持配置为64或128
- blockShapeOptional如果传入,则必须包含至少两个元素[blockShapeX, blockShapeY],且值必须大于0,blockShapeY必须为128的倍数。
- blockSparseMaskOptional当前必须传入,且shape必须为[batch, headNum, ceilDiv(maxQS, blockShapeX), ceilDiv(maxKVS, blockShapeY)]。
- attentionMaskOptional当前只支持传入nullptr。
- actualSeqLengthsOptional在qInputLayout为“TND”时必选;actualSeqLengthsKvOptional在kvInputLayout为“TND”时必选。
- actualSeqLengthsOptional与actualSeqLengthsKvOptional当前必须同时配置或同时不配置,仅配置其中之一的行为将被算子拦截。
- blockTableOptional当前只支持传入nullptr,表示不开启PagedAttention特性。
- innerPrecise仅支持配置4,表示混合精度运算,在性能与精度上取得一个折中。
- softmaxLseFlag仅支持配置0或1,分别表示不开启/开启softmaxLse输出。
- qSeqlen和kvSeqlen不需要被blockShape整除,支持非对齐场景,实际分块数通过向上取整计算。
- 输入query的headNum为N1,输入key和value的headNum为N2,则N1 >= N2 && N1 % N2 == 0。
- maskType当前只支持输入0,表示不加mask。
- blockSize当前只支持输入0,表示不支持paged cache。
- preTokens和nextTokens当前只支持输入2147483647,表示当前token的前后所有token都参与attention运算,即不支持滑窗attention。
- FP8相关约束(新增):
- 仅[object Object]Ascend 950PR/Ascend 950DT[object Object]支持。
- 当query、key、value中任意一个数据类型为FLOAT8_E4M3FN时,query、key、value必须同时为FLOAT8_E4M3FN数据类型。
- 使用FP8输入时,必须提供对应的量化缩放因子输入qDequantScale、kDequantScale、vDequantScale。
- 量化缩放因子的数据类型必须为FLOAT32。
- qDequantScale的shape必须为(Batch, HeadNum, CeilDiv(maxQSeqLength, 128), 1)。
- kDequantScale和vDequantScale的shape必须一致,为(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 256), 1)或(Batch, KVHeadNum, CeilDiv(maxKVSeqLength, 512), 1)。
- 当query、key、value中任意一个数据类型不为FLOAT8_E4M3FN时,qDequantScale、kDequantScale、vDequantScale必须传入nullptr。
- blockShapeOptional必须传入。
- q和kv的量化块大小必须与blockShapeOptional的两个元素大小分别保持一致。
[object Object]