aclnnSparseLightningIndexerGradKLLoss

产品支持情况

  • Ascend 950PR/Ascend 950DT:支持
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:支持
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:支持
  • Atlas 200I/500 A2 推理产品:不支持
  • Atlas 推理系列产品:不支持
  • Atlas 训练系列产品:不支持

功能说明

  • 接口功能:SparselightningIndexerGradKlLoss算子是LightningIndexer的反向算子,再额外融合了Loss计算功能。LightningIndexer算子将QueryToken和KeyToken之间的最高内在联系的TopK个筛选出来,存放在SparseIndices中,从而减少长序列场景下Attention的计算量,加速长序列的网络的推理和训练的性能。

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

    It,:=Wt,:@ReLU(qt,:@(K:t,:)T)I_{t,:}=W_{t,:}@ReLU(q_{t,:}@(K_{:t,:})^T)

    其中,WW是第tt个token对应的weights,qq是第tt个token对应的GG个query头合轴后的矩阵,KKttKK矩阵。

    LightningIndexer会单独训练,对应的loss function为:

    L(I)=∑tDKL(pt,:∣∣Softmax(It,:))L(I){=}\sum_tD_{KL}(p_{t,:}||Softmax(I_{t,:}))

    其中,pp是target distribution,通过对main attention score进行所有的head的求和,然后把求和结果沿着上下文方向进行L1正则化得到。DKLD_{KL}为KL散度,其表达式为:

    DKL(a∣∣b)=∑iailog(aibi)D_{KL}(a||b){=}\sum_ia_i\mathrm{log}{\left(\frac{a_i}{b_i}\right)}

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

    dIt,:=Softmax(It,:)−pt,:dI\mathop{{}}\nolimits_{{t,:}}=Softmax \left( I\mathop{{}}\nolimits_{{t,:}} \left) -p\mathop{{}}\nolimits_{{t,:}}\right. \right.

    利用链式法则可以进行weights,query和key矩阵的梯度计算:

    dWt,:=dIt,:@(ReLU(St,:))TdW\mathop{{}}\nolimits_{{t,:}}=dI\mathop{{}}\nolimits_{{t,:}}\text{@} \left( ReLU \left( S\mathop{{}}\nolimits_{{t,:}} \left) \left) \mathop{{}}\nolimits^{{T}}\right. \right. \right. \right.

    dqt,:=dSt,:@K:t,:d\mathop{{q}}\nolimits_{{t,:}}=dS\mathop{{}}\nolimits_{{t,:}}@K\mathop{{}}\nolimits_{{:t,:}}

    dK:t,:=(dSt,:)T@q:t,:dK\mathop{{}}\nolimits_{{:t,:}}= \left( dS\mathop{{}}\nolimits_{{t,:}} \left) \mathop{{}}\nolimits^{{T}}@q\mathop{{}}\nolimits_{{:t,:}}\right. \right.

    其中,S为QK矩阵softmax的结果。

函数原型

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

aclnnStatus aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize(
    const aclTensor     *query,
    const aclTensor     *key,
    const aclTensor     *queryIndex,
    const aclTensor     *keyIndex,
    const aclTensor     *weights,
    const aclTensor     *sparseIndices,
    const aclTensor     *softmaxMax,
    const aclTensor     *softmaxSum,
    const aclTensor     *queryRope,
    const aclTensor     *keyRope,
    const aclIntArray   *actualSeqLengthsQuery,
    const aclIntArray   *actualSeqLengthsKey,
    double               scaleValue,
    char                *layout,
    int64_t              sparseMode,
    int64_t              pre_tokens,
    int64_t              next_tokens,
    bool                 deterministic,
    const aclTensor     *dQueryIndex,
    const aclTensor     *dKeyIndex,
    const aclTensor     *dWeights,
    const aclTensor     *loss,
    uint64_t            *workspaceSize,
    aclOpExecutor       **executor)
aclnnStatus aclnnSparseLightningIndexerGradKLLoss(
    void             *workspace,
    uint64_t          workspaceSize,
    aclOpExecutor    *executor,
    aclrtStream stream)

aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize

  • 参数说明:

    参数名 输入/输出 描述 使用说明 数据类型 数据格式 维度(shape) 非连续Tensor
    query 输入 attention结构的输入Q。
    • 数据类型与key/queryIndex/keyIndex保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S1,N1,DQuery)、(T1,N1,DQuery)
    key 输入 attention结构的输入K。
    • 数据类型与query/queryIndex/keyIndex保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S2,N2,DQuery)、(T2,N2,DQuery)
    queryIndex 输入 lightingIndexer结构的输入queryIndex。
    • 数据类型与query/key/keyIndex保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S1,Nidx1,DQueryIndex)、(T1,Nidx1,DQueryIndex)
    keyIndex 输入 lightingIndexer结构的输入keyIndex。
    • 数据类型与query/key/queryIndex保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S2,Nidx2,DQueryIndex)、(T2,Nidx2,DQueryIndex)
    weights 输入 权重。 不支持空Tensor。 FLOAT16、BFLOAT16、FLOAT32 ND (B,S1,Nidx1)、(T1,Nidx1)
    sparseIndices 输入 topk_index,用来选择每个query对应的key和value。 不支持空Tensor。 INT32 ND (B,S1,Nidx2,K)、(T1,Nidx2,K)
    softmaxMax 输入 注意力正向计算的中间输出。 不支持空Tensor。 FLOAT32 ND (B,N2,S1,G)、(N2,T1,G)
    softmaxSum 输入 注意力正向计算的中间输出。 不支持空Tensor。 FLOAT32 ND (B,N2,S1,G)、(N2,T1,G)
    queryRope 输入 MLA rope部分:Query位置编码的输出。
    • 与query的layout维度保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S1,N1,DRope)、(T1,N1,DRope)
    keyRope 输入 MLA rope部分:Key位置编码的输出。
    • 与key的layout维度保持一致。
    • 不支持空Tensor。
    FLOAT16、BFLOAT16 ND (B,S2,N2,DRope)、(T2,N2,DRope)
    actualSeqLengthsQuery 输入 每个Batch中,Query的有效token数。
    • 值依赖。
    • 长度与B保持一致。
    • 累加和与T1保持一致。
    • 不支持空Tensor。
    INT64 ND (B,)
    actualSeqLengthsKey 输入 每个Batch中,Key的有效token数。
    • 值依赖。
    • 长度与B保持一致。
    • 累加和T2保持一致。
    • 不支持空Tensor。
    INT64 ND (B,)
    scaleValue 输入 缩放系数。 建议值:公式中d开根号的倒数。 - - - -
    layout 输入 layout格式。 仅支持BSND和TND格式。 STRING - - -
    sparseMode 输入 sparse的模式。
    • 表示sparse的模式。sparse不同模式的详细说明请参见约束说明
    • 仅支持模式3。
    INT64 - - -
    deterministic 输入 确定性计算。 优先使用整网确定性配置,该参数不产生任何效果。 BOOL - - -
    dQueryIndex 输出 QueryIndex的梯度。 - FLOAT16、BFLOAT16 ND (B,S1,Nidx1,DQueryIndex)、(T1,Nidx1,DQueryIndex) x
    dKeyIndex 输出 KeyIndex的梯度。 - FLOAT16、BFLOAT16 ND (B,S2,Nidx2,DQueryIndex)、(T2,Nidx2,DQueryIndex) x
    dWeights 输出 Weights的梯度。 - FLOAT16、BFLOAT16、FLOAT32 ND (B,S1,Nidx1)、(T1,Nidx1) x
    loss 输出 损失函数值。 - FLOAT32 ND (1,) x
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:
    • T1支持大于等于actualSeqLengthsQuery的累加和,T2支持大于等于actualSeqLengthsKey的累加和。
  • 返回值:

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

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

    返回值 错误码 描述
    ACLNN_ERR_PARAM_NULLPTR 161001 必选参数或者输出是空指针。
    ACLNN_ERR_PARAM_INVALID 161002 query、key、queryIndex、keyIndex、weights、sparseIndices、softmaxMax等输入变量的数据类型和数据格式不在支持的范围内。
    ACLNN_ERR_RUNTIME_ERROR 361001 API内存调用npu runtime的接口异常。

aclnnSparseLightningIndexerGradKLLoss

  • 参数说明:

    参数名 输入/输出 描述
    workspace 输入 在Device侧申请的workspace内存地址。
    workspaceSize 输入 在Device侧申请的workspace大小,由第一段接口aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize获取。
    executor 输入 op执行器,包含了算子计算流程。
    stream 输入 指定执行任务的Stream。
  • 返回值:

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

约束说明

  • 确定性计算:

    • aclnnSparseLightningIndexerGradKLLoss默认非确定性实现,支持通过aclrtCtxSetSysParamOpt开启确定性。
  • 公共约束

    • 参数query、key、queryIndex、keyIndex的数据类型应保持一致。
    • 参数weights不为float32时,参数query、key、queryIndex、keyIndex、weights的数据类型应保持一致。
    • 入参为空的场景处理:
      • query为空Tensor:直接返回。
      • 公共约束里入参为空的场景和FAG保持一致。
    sparseMode 含义 备注
    0 defaultMask模式,如果attenmask未传入则不做mask操作,忽略preTokens和nextTokens;如果传入,则需要传入完整的attenmask矩阵,表示preTokens和nextTokens之间的部分需要计算 不支持
    1 allMask,必须传入完整的attenmask矩阵 不支持
    2 leftUpCausal模式的mask,需要传入优化后的attenmask矩阵 不支持
    3 rightDownCausal模式的mask,对应以右顶点为划分的下三角场景,需要传入优化后的attenmask矩阵 支持
    4 band模式的mask,需要传入优化后的attenmask矩阵 不支持
    5 prefix 不支持
    6 global 不支持
    7 dilated 不支持
    8 block_local 不支持
  • 规格约束

    规格项 规格 规格说明
    B 支持1~256 -
    S1、S2 S1支持1~8K,S2支持1~512K S1、S2支持不等长;S1必须小于等于S2
    N1 32、64、128 SparseFA为MQA。
    Nidx1 8、16、32、64 SparseFA为MQA。
    N2 1 SparseFA为MQA,Nidx2=1。
    Nidx2 1 SparseFA为MQA,N2=1。
    DQuery 512 -
    DQueryIndex 128 -
    DRope 64 -
    K 1024、2048、3072、4096、5120、6144、7168、8192 -

    Ascend 950PR/Ascend 950DT:N1额外支持48,Nidx1额外支持24,二者仅允许(48,24)组合,禁止其余数值配对。

    Ascend 950PR/Ascend 950DT:B、S1、S2均支持泛化。

调用示例

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

#include <iostream>
#include <vector>
#include "acl/acl.h"
#include "aclnnop/aclnn_sparse_lightning_indexer_grad_kl_loss.h"

#define CHECK_RET(cond, return_expr) \
  do {                               \
    if (!(cond)) {                   \
      return_expr;                   \
    }                                \
  } while (0)

#define LOG_PRINT(message, ...)     \
  do {                              \
    printf(message, ##__VA_ARGS__); \
  } while (0)

int64_t GetShapeSize(const std::vector<int64_t>& shape) {
  int64_t shapeSize = 1;
  for (auto i : shape) {
    shapeSize *= i;
  }
  return shapeSize;
}

void PrintOutResult(std::vector<int64_t> &shape, void** deviceAddr) {
  auto size = GetShapeSize(shape);
  std::vector<float> resultData(size, 0);
  auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]),
                         *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return);
  for (int64_t i = 0; i < size; i++) {
    LOG_PRINT("mean result[%ld] is: %f\n", i, resultData[i]);
  }
}

void PrintOutResultFp16(std::vector<int64_t> &shape, void** deviceAddr) {
  auto size = GetShapeSize(shape);
  std::vector<aclFloat16> resultData(size, 0);
  auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]),
                         *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return);
  for (int64_t i = 0; i < size; i++) {
    LOG_PRINT("mean result[%ld] is: %f\n", i, aclFloat16ToFloat(resultData[i]));
  }
}

int Init(int32_t deviceId, aclrtContext* context, aclrtStream* stream) {
  // 固定写法,AscendCL初始化
  auto ret = aclInit(nullptr);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret);
  ret = aclrtSetDevice(deviceId);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret);
  ret = aclrtCreateContext(context, deviceId);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret);
  ret = aclrtSetCurrentContext(*context);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret);
  ret = aclrtCreateStream(stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret);
  return 0;
}

template <typename T>
int CreateAclTensor(const std::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr,
                    aclDataType dataType, aclTensor** tensor) {
  auto size = GetShapeSize(shape) * sizeof(T);
  // 调用aclrtMalloc申请device侧内存
  auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret);
  // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上
  ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret);

  // 计算连续tensor的strides
  std::vector<int64_t> strides(shape.size(), 1);
  for (int64_t i = shape.size() - 2; i >= 0; i--) {
    strides[i] = shape[i + 1] * strides[i + 1];
  }

  // 调用aclCreateTensor接口创建aclTensor
  *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
                            shape.data(), shape.size(), *deviceAddr);
  return 0;
}

int main() {
  // 1.(固定写法)device/context/stream初始化,参考AscendCL对外接口列表
  // 根据自己的实际device填写deviceId
  int32_t deviceId = 3;
  aclrtContext context;
  aclrtStream stream;
  auto ret = Init(deviceId, &context, &stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret);

  // 2. 构造输入与输出,需要根据API的接口自定义构造
  std::vector<int64_t> qShape = {1,64,512};
  std::vector<int64_t> kShape = {1,1,512};
  std::vector<int64_t> qRopeShape = {1,64,64};
  std::vector<int64_t> kRopeShape = {1,1,64};
  std::vector<int64_t> qIndexShape = {1,32,128};
  std::vector<int64_t> kIndexShape = {1,1,128};
  std::vector<int64_t> weightShape = {1,32};
  std::vector<int64_t> sparseIndicesShape = {1, 1,2048};
  std::vector<int64_t> softmaxMaxShape = {1, 1, 64};
  std::vector<int64_t> softmaxSumShape = {1, 1, 64};

  std::vector<int64_t> dQIndexShape = {1,32,128};
  std::vector<int64_t> dKIndexShape = {1,1,128};
  std::vector<int64_t> dWeightShape = {1,32};
  std::vector<int64_t> lossShape = {1};

  void* qDeviceAddr = nullptr;
  void* kDeviceAddr = nullptr;
  void* qRopeDeviceAddr = nullptr;
  void* kRopeDeviceAddr = nullptr;
  void* qIndexDeviceAddr = nullptr;
  void* kIndexDeviceAddr = nullptr;
  void* weightDeviceAddr = nullptr;
  void* sparseIndicesDeviceAddr = nullptr;
  void* softmaxMaxDeviceAddr = nullptr;
  void* softmaxSumDeviceAddr = nullptr;

  void* dQIndexDeviceAddr = nullptr;
  void* dKIndexDeviceAddr = nullptr;
  void* dWeightDeviceAddr = nullptr;
  void* lossDeviceAddr = nullptr;

  aclTensor* q = nullptr;
  aclTensor* k = nullptr;
  aclTensor* qRope = nullptr;
  aclTensor* kRope = nullptr;
  aclTensor* qIndex = nullptr;
  aclTensor* kIndex = nullptr;
  aclTensor* weight = nullptr;
  aclTensor* sparseIndices = nullptr;
  aclTensor* softmaxMax = nullptr;
  aclTensor* softmaxSum = nullptr;

  aclTensor* dQIndex = nullptr;
  aclTensor* dKIndex = nullptr;
  aclTensor* dWeight = nullptr;
  aclTensor* loss = nullptr;

  std::vector<float> qHostData(1*64*512, 1);
  std::vector<float> kHostData(1*1*512, 1);
  std::vector<float> qRopeHostData(1*64*64, 1);
  std::vector<float> kRopeHostData(1*1*64, 1);
  std::vector<float> qIndexHostData(1*32*128, 1);
  std::vector<float> kIndexHostData(1*1*128, 1);
  std::vector<float> weightHostData(1*32, 1);
  std::vector<int32_t> sparseIndicesHostData(2048, 1);
  std::vector<float> softmaxMaxHostData(1*64, 1);
  std::vector<float> softmaxSumHostData(1*64, 1);

  std::vector<float> dQIndexHostData(1*32*128, 1);
  std::vector<float> dKIndexHostData(1*1*128, 1);
  std::vector<float> dWeightHostData(1*32, 1);
  std::vector<float> lossHostData(1, 1);

  ret = CreateAclTensor(qHostData, qShape, &qDeviceAddr, aclDataType::ACL_FLOAT16, &q);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(kHostData, kShape, &kDeviceAddr, aclDataType::ACL_FLOAT16, &k);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(qRopeHostData, qRopeShape, &qRopeDeviceAddr, aclDataType::ACL_FLOAT16, &qRope);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(kRopeHostData, kRopeShape, &kRopeDeviceAddr, aclDataType::ACL_FLOAT16, &kRope);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(qIndexHostData, qIndexShape, &qIndexDeviceAddr, aclDataType::ACL_FLOAT16, &qIndex);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(kIndexHostData, kIndexShape, &kIndexDeviceAddr, aclDataType::ACL_FLOAT16, &kIndex);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(weightHostData, weightShape, &weightDeviceAddr, aclDataType::ACL_FLOAT16, &weight);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(sparseIndicesHostData, sparseIndicesShape, &sparseIndicesDeviceAddr, aclDataType::ACL_INT32, &sparseIndices);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(softmaxMaxHostData, softmaxMaxShape, &softmaxMaxDeviceAddr, aclDataType::ACL_FLOAT, &softmaxMax);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(softmaxSumHostData, softmaxSumShape, &softmaxSumDeviceAddr, aclDataType::ACL_FLOAT, &softmaxSum);
  CHECK_RET(ret == ACL_SUCCESS, return ret);

  ret = CreateAclTensor(dQIndexHostData, dQIndexShape, &dQIndexDeviceAddr, aclDataType::ACL_FLOAT16, &dQIndex);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(dKIndexHostData, dKIndexShape, &dKIndexDeviceAddr, aclDataType::ACL_FLOAT16, &dKIndex);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(dWeightHostData, dWeightShape, &dWeightDeviceAddr, aclDataType::ACL_FLOAT16, &dWeight);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(lossHostData, lossShape, &lossDeviceAddr, aclDataType::ACL_FLOAT, &loss);
  CHECK_RET(ret == ACL_SUCCESS, return ret);

  std::vector<int64_t>  acSeqQLenOp = {1};
  std::vector<int64_t>  acSeqKvLenOp = {1};
  aclIntArray* acSeqQLen = aclCreateIntArray(acSeqQLenOp.data(), acSeqQLenOp.size());
  aclIntArray* acSeqKvLen = aclCreateIntArray(acSeqKvLenOp.data(), acSeqKvLenOp.size());
  double scaleValue = 0.044194173824159216;
  int64_t preTokens = 2147483647;
  int64_t nextTokens = 2147483647;
  int64_t sparseMode = 3;
  bool deterministic = false;

  char layOut[5] = {'T', 'N', 'D', 0};

  // 3. 调用CANN算子库API,需要修改为具体的Api名称
  uint64_t workspaceSize = 0;
  aclOpExecutor* executor;

  // 调用aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize第一段接口
  ret = aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize(
            q, k, qIndex, kIndex, weight, sparseIndices, softmaxMax, softmaxSum, qRope, kRope, acSeqQLen, acSeqKvLen,
            scaleValue, layOut, sparseMode, preTokens, nextTokens, deterministic, dQIndex, dKIndex, dWeight, loss,
            &workspaceSize, &executor);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseLightningIndexerGradKLLossGetWorkspaceSize failed. ERROR: %d\n", ret);
            return ret);

  // 根据第一段接口计算出的workspaceSize申请device内存
  void* workspaceAddr = nullptr;
  if (workspaceSize > 0) {
    ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
    CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret);
  }

  // 调用aclnnSparseLightningIndexerGradKLLoss第二段接口
  ret = aclnnSparseLightningIndexerGradKLLoss(workspaceAddr, workspaceSize, executor, stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseLightningIndexerGradKLLoss failed. ERROR: %d\n", ret); return ret);

  // 4.(固定写法)同步等待任务执行结束
  ret = aclrtSynchronizeStream(stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret);

  // 5. 获取输出的值,将device侧内存上的结果拷贝至host侧,需要根据具体API的接口定义修改
  PrintOutResultFp16(dQIndexShape, &dQIndexDeviceAddr);
  PrintOutResultFp16(dKIndexShape, &dKIndexDeviceAddr);
  PrintOutResultFp16(dWeightShape, &dWeightDeviceAddr);
  PrintOutResult(lossShape, &lossDeviceAddr);

  // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改
  aclDestroyTensor(q);
  aclDestroyTensor(k);
  aclDestroyTensor(qIndex);
  aclDestroyTensor(kIndex);
  aclDestroyTensor(qRope);
  aclDestroyTensor(kRope);
  aclDestroyTensor(weight);
  aclDestroyTensor(sparseIndices);
  aclDestroyTensor(softmaxMax);
  aclDestroyTensor(softmaxSum);

  aclDestroyTensor(dQIndex);
  aclDestroyTensor(dKIndex);
  aclDestroyTensor(dWeight);
  aclDestroyTensor(loss);

  // 7. 释放device资源
  aclrtFree(qDeviceAddr);
  aclrtFree(kDeviceAddr);
  aclrtFree(qIndexDeviceAddr);
  aclrtFree(kIndexDeviceAddr);
  aclrtFree(qRopeDeviceAddr);
  aclrtFree(kRopeDeviceAddr);
  aclrtFree(weightDeviceAddr);
  aclrtFree(sparseIndicesDeviceAddr);
  aclrtFree(softmaxMaxDeviceAddr);
  aclrtFree(softmaxSumDeviceAddr);

  aclrtFree(dQIndexDeviceAddr);
  aclrtFree(dKIndexDeviceAddr);
  aclrtFree(dWeightDeviceAddr);
  aclrtFree(lossDeviceAddr);
  if (workspaceSize > 0) {
    aclrtFree(workspaceAddr);
  }
  aclrtDestroyStream(stream);
  aclrtDestroyContext(context);
  aclrtResetDevice(deviceId);
  aclFinalize();

  return 0;
}