已合并
fix nullptr check #5018
liuyun_nj创建于 5月19日
fix nullptr check #5018
已合并
共 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 | // 支持空tensor | 232 | // 支持空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()); |