已合并
fix nullptr check #5018
liuyun_nj创建于 5月19日
fix nullptr check #5018
已合并
liuyun_nj创建于 5月19日
1 个文件变更+8-8
@@ -221,6 +221,14 @@ aclnnStatus aclnnAddRmsNormGetWorkspaceSize(
221 auto uniqueExecutor = CREATE_EXECUTOR();221 auto uniqueExecutor = CREATE_EXECUTOR();
222 CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);222 CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
223 223 
224+ // 参数检查
225+ AddRmsNormACLNN::AddRmsNormInputTensor inputTensorOri = {x1, x2, gamma};
226+ AddRmsNormACLNN::AddRmsNormOutputTensor outputTensor = {yOut, rstdOut, xOut};
227+ 
228+ int64_t mode = AddRmsNormACLNN::ADD_RMS_NORM_MODE; // 0为addrmsnorm,1为preRmsNorm, 2为postRmsNorm
229+ auto ret = CheckParams(inputTensorOri, outputTensor, mode);
230+ CHECK_RET(ret == ACLNN_SUCCESS, ret);
231+ 
224 // 支持空tensor232 // 支持空tensor
225 bool anyEmptyTensor = x1->IsEmpty() || gamma->IsEmpty();233 bool anyEmptyTensor = x1->IsEmpty() || gamma->IsEmpty();
226 if (anyEmptyTensor) {234 if (anyEmptyTensor) {
@@ -230,14 +238,6 @@ aclnnStatus aclnnAddRmsNormGetWorkspaceSize(
230 return ACLNN_SUCCESS;238 return ACLNN_SUCCESS;
231 }239 }
232 240 
233- // 参数检查
234- AddRmsNormACLNN::AddRmsNormInputTensor inputTensorOri = {x1, x2, gamma};
235- AddRmsNormACLNN::AddRmsNormOutputTensor outputTensor = {yOut, rstdOut, xOut};
236- 
237- int64_t mode = AddRmsNormACLNN::ADD_RMS_NORM_MODE; // 0为addrmsnorm,1为preRmsNorm, 2为postRmsNorm
238- auto ret = CheckParams(inputTensorOri, outputTensor, mode);
239- CHECK_RET(ret == ACLNN_SUCCESS, ret);
240- 
241 // 固定写法,将输入转换成连续的tensor,可选输入不做判空校验241 // 固定写法,将输入转换成连续的tensor,可选输入不做判空校验
242 auto x1Cont = l0op::Contiguous(x1, uniqueExecutor.get());242 auto x1Cont = l0op::Contiguous(x1, uniqueExecutor.get());
243 auto x2Cont = l0op::Contiguous(x2, uniqueExecutor.get());243 auto x2Cont = l0op::Contiguous(x2, uniqueExecutor.get());