Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
bsa_select_block_mask(Block Sparse Attention 块选择掩码算子)用于 NSA(Native Sparse Attention)类稀疏注意力的块级稀疏掩码生成。原有流程为 4 步: Q/K Mean Pooling → QK Matmul → Online Softmax → Radix TopK → 输出 mask 即对 Q/K 各做一次 pooling 压缩后,在 [xBlocks, yBlocks] 细粒度块级得分矩阵上直接 TopK 选块。随着序列长度增长到 256K~6M 级别,细粒度块矩阵本身规模巨大,一次压缩的粒度已不满足需求方的稀疏控制需求。 此次需求:希望把一次 pooling 拆成两次——在 softmax 得分矩阵上再叠加一次 Mean Pooling(由新增可选输入 post_block_shape 指定二次池化窗口),TopK 在粗粒度矩阵 pooled_score[postXBlocks, postYBlocks] 上选择,直接输出粗粒度掩码 [B, N, postXBlocks, postYBlocks]。
通过二次 pooling 将 TopK 选择粒度从细粒度块提升到粗粒度块,下游注意力仅对选中的粗块计算,显著降低长序列下的计算与访存量。
接口功能:aclnnBSASelectBlockMask是BSA(BlockSparseAttention)的前置算子,负责根据Query和Key的内容动态生成blockSparseMask,使BSA的调用链从"手动提供掩码"变为"根据Q/K内容自适应选择稀疏模式"。
计算公式: 设blockShape = [blockShapeX, blockShapeY],Sq是query最大序列长度,Skv是key最大序列长度, 则压缩后块数:
Xblocks=⌈Sq/blockShapeX⌉,Yblocks=⌈Skv/blockShapeY⌉Xblocks = \lceil Sq / blockShapeX \rceil,\quad Yblocks = \lceil Skv / blockShapeY \rceil Xblocks=⌈Sq/blockShapeX⌉,Yblocks=⌈Skv/blockShapeY⌉
Step1:均值池化压缩 (Mean Pooling Compression) 当actualBlockLenQuery / actualBlockLenKey为null时(完整压缩):
q_compressed[b,n,x,d]=1blockShapeX∑i=0blockShapeX−1query[b,n,x⋅blockShapeX+i,d]q\_compressed[b, n, x, d] = \frac{1}{blockShapeX} \sum_{i=0}^{blockShapeX-1} query[b, n, x \cdot blockShapeX + i, d] q_compressed[b,n,x,d]=blockShapeX1i=0∑blockShapeX−1query[b,n,x⋅blockShapeX+i,d]
k_compressed[b,n,y,d]=1blockShapeY∑j=0blockShapeY−1key[b,n,y⋅blockShapeY+j,d]k\_compressed[b, n, y, d] = \frac{1}{blockShapeY} \sum_{j=0}^{blockShapeY-1} key[b, n, y \cdot blockShapeY + j, d] k_compressed[b,n,y,d]=blockShapeY1j=0∑blockShapeY−1key[b,n,y⋅blockShapeY+j,d]
当actualBlockLenQuery / actualBlockLenKey非null时(部分压缩),仅对每个block内前actualBlockLen个token取均值:
q_compressed[b,n,x,d]=1actualBlockLenQ[b,x]∑i=0actualBlockLenQ[b,x]−1query[b,n,x⋅blockShapeX+i,d]q\_compressed[b, n, x, d] = \frac{1}{actualBlockLenQ[b,x]} \sum_{i=0}^{actualBlockLenQ[b,x]-1} query[b, n, x \cdot blockShapeX + i, d] q_compressed[b,n,x,d]=actualBlockLenQ[b,x]1i=0∑actualBlockLenQ[b,x]−1query[b,n,x⋅blockShapeX+i,d]
k_compressed[b,n,y,d]=1actualBlockLenK[b,y]∑j=0actualBlockLenK[b,y]−1key[b,n,y⋅blockShapeY+j,d]k\_compressed[b, n, y, d] = \frac{1}{actualBlockLenK[b,y]} \sum_{j=0}^{actualBlockLenK[b,y]-1} key[b, n, y \cdot blockShapeY + j, d] k_compressed[b,n,y,d]=actualBlockLenK[b,y]1j=0∑actualBlockLenK[b,y]−1key[b,n,y⋅blockShapeY+j,d]
Step2a:QK Matmul
score[b,n,x,y]=scale⋅∑d=0D−1q_compressed[b,n,x,d]⋅k_compressed[b,n,y,d]score[b, n, x, y] = scale \cdot \sum_{d=0}^{D-1} q\_compressed[b, n, x, d] \cdot k\_compressed[b, n, y, d] score[b,n,x,y]=scale⋅d=0∑D−1q_compressed[b,n,x,d]⋅k_compressed[b,n,y,d]
Step2b:Softmax
attn_score[b,n,x,y]=softmax(score[b,n,x,:])=exp(score[b,n,x,y]−mfinal)lfinalattn\_score[b, n, x, y] = softmax(score[b, n, x, :]) = \frac{\exp(score[b, n, x, y] - m_{final})}{l_{final}} attn_score[b,n,x,y]=softmax(score[b,n,x,:])=lfinalexp(score[b,n,x,y]−mfinal)
Step2c:二次池化压缩 (Post-Softmax Mean Pooling) 当postBlockShape非null时,对attn_score做二次均值池化,生成粗粒度pooled_score:
postXBlocks=⌈Xblocks/postBlockShapeX⌉,postYBlocks=⌈Yblocks/postBlockShapeY⌉postXBlocks = \lceil Xblocks / postBlockShapeX \rceil,\quad postYBlocks = \lceil Yblocks / postBlockShapeY \rceil postXBlocks=⌈Xblocks/postBlockShapeX⌉,postYBlocks=⌈Yblocks/postBlockShapeY⌉
pooled_score[b,n,px,py]=1∣Rpx,py∣∑x∈Rx(px)∑y∈Ry(py)attn_score[b,n,x,y]pooled\_score[b, n, px, py] = \frac{1}{|R_{px,py}|} \sum_{x \in R_x(px)} \sum_{y \in R_y(py)} attn\_score[b, n, x, y] pooled_score[b,n,px,py]=∣Rpx,py∣1x∈Rx(px)∑y∈Ry(py)∑attn_score[b,n,x,y]
当postBlockShape为null时,跳过此步骤,pooled_score = attn_score。 Step3:TopK选择生成索引
topk_value=round(sparsity×postXBlocks×postYBlocks)topk\_value = \text{round}(sparsity \times postXBlocks \times postYBlocks) topk_value=round(sparsity×postXBlocks×postYBlocks)
indices=TopK(pooled_score[b,n,px,py], topK_value)\mathcal{indices}= \text{TopK}\left(pooled\_score[b, n, px, py],\; topK\_value\right) indices=TopK(pooled_score[b,n,px,py],topK_value)
当postBlockShape为null时,postXBlocks=Xblocks、postYBlocks=Yblocks、pooled_score=attn_score,等价于直接在attn_score上做TopK。 其中indices为attn_score[b, n, x, y] 中topk_value个最大值对应的索引集合。 Step4:生成BlockSparseMask 当postBlockShape非null时,输出直接为粗粒度mask(shape为[B, N, postXBlocks, postYBlocks],二次pooling拆分语义,无需展开):
blockSparseMaskOut[b,n,px,py]={1(px,py)∈indices0(px,py)∉indicesblockSparseMaskOut[b, n, px, py] = \begin{cases} 1 & (px, py) \in \mathcal{indices} \\ 0 & (px, py) \notin \mathcal{indices} \end{cases} blockSparseMaskOut[b,n,px,py]={10(px,py)∈indices(px,py)∈/indices
当postBlockShape为null时,输出为细粒度mask(shape为[B, N, Xblocks, Yblocks]),直接逐元素生成:
blockSparseMaskOut[b,n,x,y]={1(b,n,x,y)∈indices0(b,n,x,y)∉indicesblockSparseMaskOut[b, n, x, y] = \begin{cases} 1 & (b, n, x, y) \in \mathcal{indices} \\ 0 & (b, n, x, y) \notin \mathcal{indices} \end{cases} blockSparseMaskOut[b,n,x,y]={10(b,n,x,y)∈indices(b,n,x,y)∈/indices
/assign @dai_lin
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
bsa_select_block_mask(Block Sparse Attention 块选择掩码算子)用于 NSA(Native Sparse Attention)类稀疏注意力的块级稀疏掩码生成。原有流程为 4 步:
Q/K Mean Pooling → QK Matmul → Online Softmax → Radix TopK → 输出 mask
即对 Q/K 各做一次 pooling 压缩后,在 [xBlocks, yBlocks] 细粒度块级得分矩阵上直接 TopK 选块。随着序列长度增长到 256K~6M 级别,细粒度块矩阵本身规模巨大,一次压缩的粒度已不满足需求方的稀疏控制需求。
此次需求:希望把一次 pooling 拆成两次——在 softmax 得分矩阵上再叠加一次 Mean Pooling(由新增可选输入 post_block_shape 指定二次池化窗口),TopK 在粗粒度矩阵 pooled_score[postXBlocks, postYBlocks] 上选择,直接输出粗粒度掩码 [B, N, postXBlocks, postYBlocks]。
Benefit / Necessity (价值/作用)
通过二次 pooling 将 TopK 选择粒度从细粒度块提升到粗粒度块,下游注意力仅对选中的粗块计算,显著降低长序列下的计算与访存量。
Design(设计方案)
接口功能:aclnnBSASelectBlockMask是BSA(BlockSparseAttention)的前置算子,负责根据Query和Key的内容动态生成blockSparseMask,使BSA的调用链从"手动提供掩码"变为"根据Q/K内容自适应选择稀疏模式"。
计算公式:
设blockShape = [blockShapeX, blockShapeY],Sq是query最大序列长度,Skv是key最大序列长度, 则压缩后块数:
Xblocks=⌈Sq/blockShapeX⌉,Yblocks=⌈Skv/blockShapeY⌉
Step1:均值池化压缩 (Mean Pooling Compression)
当actualBlockLenQuery / actualBlockLenKey为null时(完整压缩):
q_compressed[b,n,x,d]=blockShapeX1i=0∑blockShapeX−1query[b,n,x⋅blockShapeX+i,d]
k_compressed[b,n,y,d]=blockShapeY1j=0∑blockShapeY−1key[b,n,y⋅blockShapeY+j,d]
当actualBlockLenQuery / actualBlockLenKey非null时(部分压缩),仅对每个block内前actualBlockLen个token取均值:
q_compressed[b,n,x,d]=actualBlockLenQ[b,x]1i=0∑actualBlockLenQ[b,x]−1query[b,n,x⋅blockShapeX+i,d]
k_compressed[b,n,y,d]=actualBlockLenK[b,y]1j=0∑actualBlockLenK[b,y]−1key[b,n,y⋅blockShapeY+j,d]
Step2a:QK Matmul
score[b,n,x,y]=scale⋅d=0∑D−1q_compressed[b,n,x,d]⋅k_compressed[b,n,y,d]
Step2b:Softmax
attn_score[b,n,x,y]=softmax(score[b,n,x,:])=lfinalexp(score[b,n,x,y]−mfinal)
Step2c:二次池化压缩 (Post-Softmax Mean Pooling)
当postBlockShape非null时,对attn_score做二次均值池化,生成粗粒度pooled_score:
postXBlocks=⌈Xblocks/postBlockShapeX⌉,postYBlocks=⌈Yblocks/postBlockShapeY⌉
pooled_score[b,n,px,py]=∣Rpx,py∣1x∈Rx(px)∑y∈Ry(py)∑attn_score[b,n,x,y]
当postBlockShape为null时,跳过此步骤,pooled_score = attn_score。
Step3:TopK选择生成索引
topk_value=round(sparsity×postXBlocks×postYBlocks)
indices=TopK(pooled_score[b,n,px,py],topK_value)
当postBlockShape为null时,postXBlocks=Xblocks、postYBlocks=Yblocks、pooled_score=attn_score,等价于直接在attn_score上做TopK。
其中indices为attn_score[b, n, x, y] 中topk_value个最大值对应的索引集合。
Step4:生成BlockSparseMask
当postBlockShape非null时,输出直接为粗粒度mask(shape为[B, N, postXBlocks, postYBlocks],二次pooling拆分语义,无需展开):
blockSparseMaskOut[b,n,px,py]={10(px,py)∈indices(px,py)∈/indices
当postBlockShape为null时,输出为细粒度mask(shape为[B, N, Xblocks, Yblocks]),直接逐元素生成:
blockSparseMaskOut[b,n,x,y]={10(b,n,x,y)∈indices(b,n,x,y)∈/indices