接口功能: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可表示为:
其中,和分别对应本接口的q和k,对应本接口的w,由sparseIndices给出。Indexer分支的softmax输出为:
本接口将写出到softmaxOut。目标分布由attnSoftmaxL1Norm输入提供,等价于旧版kernel内部由main attention score经head求和和L1归一化得到的结果。若后续继续计算KL Loss,其形式与旧版保持一致:
通过求导可得Loss的梯度表达式:
利用链式法则可以进行w、q和k矩阵的梯度计算:
dK写回时会按照sparseIndices指向的key位置做scatter-add,无效top-k位置不参与计算。
每个算子分为,必须先调用“aclnnSparseLightningIndexerKLLossGradGetWorkspaceSize”接口获取入参并根据计算流程计算所需workspace大小,再调用“aclnnSparseLightningIndexerKLLossGrad”接口执行计算。
参数说明:
[object Object][object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:暂不支持seqUsedQOptional、seqUsedKOptional字段。
返回值:
第一段接口完成入参校验,出现以下场景时报错:
[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。
- 参数B的支持情况: