开发者
下载
[object Object][object Object][object Object]undefined
[object Object]
  • 接口功能:SparseLightningIndexerKLLossGrad算子是LightningIndexer KL Loss的反向计算算子。该接口接收Lightning Indexer分支的q、k、w、sparseIndices,以及由主Attention分支预先计算得到的attnSoftmaxL1Norm,计算并输出dq、dk、dw和softmaxOut。与SparseLightningIndexerGradKLLoss相比,本接口不再在kernel内部重算主Attention的[object Object],也不再输出loss;主Attention的目标分布由attnSoftmaxL1Norm输入提供,softmaxOut可用于后续loss计算。

  • 计算公式: 用于取Top-k的value的Indexer logits可表示为:

    St,:=qt,:@Ktopk(t),:TS_{t,:}=q_{t,:}@K_{\operatorname{topk}(t),:}^{T} It,:=Wt,:@ReLU(St,:)I_{t,:}=W_{t,:}@\mathrm{ReLU}(S_{t,:})

    其中,qqKK分别对应本接口的q和k,WW对应本接口的w,topk(t)\operatorname{topk}(t)由sparseIndices给出。Indexer分支的softmax输出为:

    yt,:=Softmax(It,:)y_{t,:}=\operatorname{Softmax}(I_{t,:})

    本接口将yy写出到softmaxOut。目标分布pp由attnSoftmaxL1Norm输入提供,等价于旧版kernel内部由main attention score经head求和和L1归一化得到的结果。若后续继续计算KL Loss,其形式与旧版保持一致:

    L(I)=tDKL(pt,:Softmax(It,:))L(I){=}\sum_tD_{KL}(p_{t,:}||\operatorname{Softmax}(I_{t,:})) DKL(ab)=iailog(aibi)D_{KL}(a||b){=}\sum_ia_i\mathrm{log}{\left(\frac{a_i}{b_i}\right)}

    通过求导可得Loss的梯度表达式:

    dIt,:=Softmax(It,:)pt,:dI_{t,:}=\operatorname{Softmax}(I_{t,:})-p_{t,:}

    利用链式法则可以进行w、q和k矩阵的梯度计算:

    dWt,:=dIt,:@(ReLU(St,:))TdW_{t,:}=dI_{t,:}\text{@}\left(\mathrm{ReLU}(S_{t,:})\right)^{T} dqt,:=dSt,:@Ktopk(t),:dq_{t,:}=dS_{t,:}@K_{\operatorname{topk}(t),:} dKtopk(t),:=(dSt,:)T@qt,:dK_{\operatorname{topk}(t),:}=\left(dS_{t,:}\right)^{T}@q_{t,:}

    dK写回时会按照sparseIndices指向的key位置做scatter-add,无效top-k位置不参与计算。

[object Object]

每个算子分为,必须先调用“aclnnSparseLightningIndexerKLLossGradGetWorkspaceSize”接口获取入参并根据计算流程计算所需workspace大小,再调用“aclnnSparseLightningIndexerKLLossGrad”接口执行计算。

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

    [object Object]
  • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:暂不支持seqUsedQOptional、seqUsedKOptional字段。

  • 返回值:

    返回aclnnStatus状态码,具体参见

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

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

    [object Object]
  • 返回值:

    返回aclnnStatus状态码,具体参见

[object Object]
  • 确定性计算:

    • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:aclnnSparseLightningIndexerKLLossGrad默认非确定性实现,不支持通过aclrtCtxSetSysParamOpt开启确定性。
    • [object Object]Ascend 950PR/Ascend 950DT[object Object]:aclnnSparseLightningIndexerKLLossGrad默认非确定性实现,支持通过aclrtCtxSetSysParamOpt开启确定性。
  • 公共约束:

    • 参数q、k、dq、dk的数据类型应保持一致,支持FLOAT16和BFLOAT16。

    • 参数w、dw、attnSoftmaxL1Norm、softmaxOut的数据类型应为FLOAT32。

    • 参数sparseIndices、cuSeqLensQOptional、cuSeqLensKOptional、seqUsedQOptional、seqUsedKOptional、cmpResidualKOptional、metadataOptional的数据类型应为INT32。

    • layoutQ和layoutK当前支持BSND和TND。

    • 当layoutQ为TND时,需要传入cuSeqLensQOptional;当layoutK为TND时,需要传入cuSeqLensKOptional。

    • sparseIndices中有效位置必须位于当前batch的key序列范围内;无效位置使用-1填充。

    • attnSoftmaxL1Norm的无效top-k位置建议置零,并与sparseIndices的有效位置保持一致。

      [object Object]
  • 规格约束:

    [object Object]
    • 参数B的支持情况:
      • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:B支持1~256。
      • [object Object]Ascend 950PR/Ascend 950DT[object Object]:B>0。
    • 参数S1、S2的支持情况:
      • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:S1支持1~8K,S2支持1~512K。
      • [object Object]Ascend 950PR/Ascend 950DT[object Object]:S1>0,S2>0。
    • 参数N1的支持情况:
      • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:N1支持8、16、32、64。
      • [object Object]Ascend 950PR/Ascend 950DT[object Object]:N1支持1~128。
    • 参数K的支持情况:
      • [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:K支持512、1024、2048、4096、8192。
      • [object Object]Ascend 950PR/Ascend 950DT[object Object]:K支持1~2048。
[object Object]

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

[object Object]