开发者
下载
[object Object][object Object][object Object]undefined
[object Object]
  • 接口功能:

    [object Object]算子实现基于共享KV(Key=Value)的稀疏注意力计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。该算子适用于大语言模型训练、推理场景,通过滑动窗口和KV压缩机制大幅降低长序列注意力计算的开销。调用时需要使用[object Object]生成的任务列表[object Object]

    该算子不建议单独使用,建议与aclnnSparseFlashMla算子配合使用,形成完整的工作流。

    典型调用流程如下:

    1. 准备[object Object][object Object][object Object]、序列长度、[object Object][object Object]等输入。
    2. 调用[object Object]生成[object Object]
    3. 调用[object Object],将上一步得到的[object Object]传入主算子。
  • 计算公式:

    O=softmax(QK~Tsoftmax_scale)V~O = \text{softmax}(Q \cdot \tilde{K}^T \cdot \text{softmax\_scale}) \cdot \tilde{V}

    其中K~=V~\tilde{K} = \tilde{V}(共享KV),K~\tilde{K}由滑动窗口内的原始KV和因果边界内的压缩KV拼接而成,具体参与计算的KV范围由模板模式和mask参数决定:

    • 滑动窗口部分(oriKv):对第iS1i_{S1}个Query token,其因果对角线位置为ori_threshold=S2actS1act+iS1+1\text{ori\_threshold} = S2_{act} - S1_{act} + i_{S1} + 1,窗口范围为[max(ori_thresholdori_win_left1,0),ori_threshold+ori_win_right)[\max(\text{ori\_threshold} - \text{ori\_win\_left} - 1, 0), \text{ori\_threshold} + \text{ori\_win\_right})

    • 压缩KV部分(cmpKv):因果边界阈值为cmp_threshold=ori_thresholdcmp_ratio\text{cmp\_threshold} = \lfloor \frac{\text{ori\_threshold}}{\text{cmp\_ratio}} \rfloor。HCA场景取[0,cmp_threshold)[0, \text{cmp\_threshold})内的连续压缩KV;CSA场景通过TopK索引从压缩KV中按需收集,仅保留begin_idx<cmp_threshold\text{begin\_idx} < \text{cmp\_threshold}的块。

    注意力计算采用Online Softmax(Flash Attention V2),S2方向按512分块循环,sinks作为每行softmax的初始最大值:

    row_max(0)=sinks[g],row_sum(0)=1.0,O(0)=0\text{row\_max}^{(0)} = \text{sinks}[g], \quad \text{row\_sum}^{(0)} = 1.0, \quad O^{(0)} = 0 S(t)=QKtile(t)Tsoftmax_scaleS^{(t)} = Q \cdot K_{tile}^{(t)T} \cdot \text{softmax\_scale} row_max(t+1)=max(row_max(t),max(S(t),dim=1))\text{row\_max}^{(t+1)} = \max(\text{row\_max}^{(t)}, \max(S^{(t)}, \text{dim}=-1)) row_sum(t+1)=exp(row_max(t)row_max(t+1))row_sum(t)+exp(S(t)row_max(t+1))\text{row\_sum}^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot \text{row\_sum}^{(t)} + \sum \exp(S^{(t)} - \text{row\_max}^{(t+1)}) O(t+1)=exp(row_max(t)row_max(t+1))O(t)+exp(S(t)row_max(t+1))Vtile(t)O^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot O^{(t)} + \exp(S^{(t)} - \text{row\_max}^{(t+1)}) \cdot V_{tile}^{(t)} Ofinal=O(Ts2)/row_sum(Ts2)O_{final} = O^{(T_{s2})} / \text{row\_sum}^{(T_{s2})}
  • 符号说明

    [object Object]undefined
[object Object]

每个算子分为,必须先调用[object Object]接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用[object Object]执行实际计算。

[object Object]
[object Object]
[object Object]
  • 参数说明

    [object Object]
    • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:N1/N2支持1、2、4、8、16、32、64、128;cmp_ratio在SWA场景保持默认值1,CSA支持传入4,HCA支持传入128;block_size取值为16的倍数,最大支持1024;ori_sparse_indices当前暂不支持,cmp_sparse_indices的最后一维K2当前支持512或1024。
    • [object Object]Ascend 950PR/Ascend 950DT[object Object]:N1/N2支持2、4、8、16、32、64、128,不支持1。
  • 返回值

    aclnnStatus:返回状态码,具体参见

    第一段接口完成入参校验,出现以下场景时报错:

    • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:

      [object Object]
    • [object Object]Ascend 950PR/Ascend 950DT[object Object]:

      [object Object]
[object Object]
  • 参数说明

    [object Object]
  • 返回值

    返回aclnnStatus状态码,具体参见

[object Object]
  • 确定性计算

    • aclnnSparseFlashMla默认采用确定性实现,相同输入多次调用结果一致。
  • 使用约束

    • 资料支持范围内暂不支持对[object Object]进行稀疏计算,设置[object Object]无效。
    • [object Object][object Object]等预留输入可传入nullptr或空Tensor外,其余已传入Tensor不支持为空。
    • [object Object]参数必须传入,由[object Object]算子生成,shape固定为(1024,)。
    • [object Object]为主算子和[object Object]的可选入参;传入后用于按[object Object]恢复cmp侧mask使用的压缩前长度。
  • 三种Attention场景输入要求

    [object Object]undefined
  • Layout约束

    • [object Object][object Object]组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下[object Object][object Object]必须一致。
    • [object Object]为TND时,[object Object]必须传入。
    • [object Object]为PA_BBND时,[object Object]必须传入,[object Object]必须传入。BSND场景可选传入[object Object]覆盖每个batch的oriKv有效长度;TND场景使用[object Object]表达oriKv序列边界。
    • [object Object]为TND时,[object Object]必须传入。
    • [object Object]为TND且存在[object Object]时,[object Object]必须传入。
    • [object Object]为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。
[object Object]

调用示例代码如下,仅供参考,具体编译和执行过程请参考

[object Object]