已关闭
[Requirement|需求建议]: bsa_select_block_mask算子新增二次pooling功能 #4634
dailin创建于  20 天前关闭于  19 天前
dailin
dailin
20 天前 创建

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/blockShapeYXblocks = \lceil Sq / blockShapeX \rceil,\quad Yblocks = \lceil Skv / blockShapeY \rceil

    Step1:均值池化压缩 (Mean Pooling Compression)
    当actualBlockLenQuery / actualBlockLenKey为null时(完整压缩):

    q_compressed[b,n,x,d]=1blockShapeXi=0blockShapeX1query[b,n,xblockShapeX+i,d]q\_compressed[b, n, x, d] = \frac{1}{blockShapeX} \sum_{i=0}^{blockShapeX-1} query[b, n, x \cdot blockShapeX + i, d]

    k_compressed[b,n,y,d]=1blockShapeYj=0blockShapeY1key[b,n,yblockShapeY+j,d]k\_compressed[b, n, y, d] = \frac{1}{blockShapeY} \sum_{j=0}^{blockShapeY-1} key[b, n, y \cdot 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,xblockShapeX+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]

    k_compressed[b,n,y,d]=1actualBlockLenK[b,y]j=0actualBlockLenK[b,y]1key[b,n,yblockShapeY+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]

    Step2a:QK Matmul

    score[b,n,x,y]=scaled=0D1q_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]

    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}}

    Step2c:二次池化压缩 (Post-Softmax Mean Pooling)
    当postBlockShape非null时,对attn_score做二次均值池化,生成粗粒度pooled_score:

    postXBlocks=Xblocks/postBlockShapeX,postYBlocks=Yblocks/postBlockShapeYpostXBlocks = \lceil Xblocks / postBlockShapeX \rceil,\quad postYBlocks = \lceil Yblocks / postBlockShapeY \rceil

    pooled_score[b,n,px,py]=1Rpx,pyxRx(px)yRy(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]

    当postBlockShape为null时,跳过此步骤,pooled_score = attn_score。
    Step3:TopK选择生成索引

    topk_value=round(sparsity×postXBlocks×postYBlocks)topk\_value = \text{round}(sparsity \times postXBlocks \times 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)

    当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}

    当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}

likedislike
weihao18成员
20 天前 评论:

/assign @dai_lin

likedislike
CANN-robotCANN-robot成员
20 天前 将 tramp-ll 设为负责人
CANN-robotCANN-robot成员
20 天前 将 dai_lin 设为负责人,移除负责人 tramp-ll
CANN-robotCANN-robot成员
19 天前 关闭了 issue
CANN-robotCANN-robot成员
19 天前 添加了label:resolved