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头合轴后的矩阵,KK为tt行KK矩阵。
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;
}