已合并
refactor: 简化 convolution_backward 空指针防御性校验并清理冗余空行 #9309
Nice try创建于 8月27日
refactor: 简化 convolution_backward 空指针防御性校验并清理冗余空行 #9309
已合并
共 2 个文件变更+27-36
| @@ -339,7 +339,6 @@ void Conv3DDWV2BasicBlockTiling::AdjustSingleNForStreamK() | |||
| 339 | uint64_t maxStreamKDim = mmInfo_.kValue / blockTiling_.blockBaseK; | 339 | uint64_t maxStreamKDim = mmInfo_.kValue / blockTiling_.blockBaseK; |
| 340 | uint64_t batchDoutDim = !context_->GetDeterministic() ? 1 : static_cast<uint64_t>(runInfo_.batch) * runInfo_.dout; | 340 | uint64_t batchDoutDim = !context_->GetDeterministic() ? 1 : static_cast<uint64_t>(runInfo_.batch) * runInfo_.dout; |
| 341 | maxStreamKDim = std::max(maxStreamKDim, batchDoutDim); | 341 | maxStreamKDim = std::max(maxStreamKDim, batchDoutDim); |
| 342 | - | ||
| 343 | if (maxStreamKDim <= STREAM_K_DIM_MIN || tailCnt + STREAM_K_TAIL_TOLERANCE >= targetCoreNum) { | 342 | if (maxStreamKDim <= STREAM_K_DIM_MIN || tailCnt + STREAM_K_TAIL_TOLERANCE >= targetCoreNum) { |
| 344 | return; | 343 | return; |
| 345 | } | 344 | } |
| @@ -3438,23 +3438,19 @@ aclnnStatus aclnnConvolutionBackwardGetWorkspaceSize( | |||
| 3438 | Ops::NN::Conv::ConvolutionBackwardChecker convolutionBackwardChecker(inputTensor, outputTensor, params, npuArch); | 3438 | Ops::NN::Conv::ConvolutionBackwardChecker convolutionBackwardChecker(inputTensor, outputTensor, params, npuArch); |
| 3439 | auto ret = convolutionBackwardChecker.CheckParams(); | 3439 | auto ret = convolutionBackwardChecker.CheckParams(); |
| 3440 | CHECK_RET(ret == ACLNN_SUCCESS, ret); | 3440 | CHECK_RET(ret == ACLNN_SUCCESS, ret); |
| 3441 | + CHECK_RET(input != nullptr && weight != nullptr && gradOutput != nullptr, ACLNN_ERR_PARAM_NULLPTR); | ||
| 3441 | 3442 | ||
| 3442 | - if (gradOutput != nullptr) { | 3443 | + auto gradOutputContiguous = l0op::Contiguous(gradOutput, uniqueExecutor.get()); |
| 3443 | - auto gradOutputContiguous = l0op::Contiguous(gradOutput, uniqueExecutor.get()); | 3444 | + CHECK_RET(gradOutputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3444 | - CHECK_RET(gradOutputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | 3445 | + inputTensor.gradOutput = gradOutputContiguous; |
| 3445 | - inputTensor.gradOutput = gradOutputContiguous; | ||
| 3446 | - } | ||
| 3447 | 3446 | ||
| 3448 | - if (input != nullptr) { | 3447 | + auto inputContiguous = l0op::Contiguous(input, uniqueExecutor.get()); |
| 3449 | - auto inputContiguous = l0op::Contiguous(input, uniqueExecutor.get()); | 3448 | + CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3450 | - CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | 3449 | + inputTensor.input = inputContiguous; |
| 3451 | - inputTensor.input = inputContiguous; | 3450 | + |
| 3452 | - } | 3451 | + auto weightContiguous = l0op::Contiguous(weight, uniqueExecutor.get()); |
| 3453 | - if (weight != nullptr) { | 3452 | + CHECK_RET(weightContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3454 | - auto weightContiguous = l0op::Contiguous(weight, uniqueExecutor.get()); | 3453 | + inputTensor.weight = weightContiguous; |
| 3455 | - CHECK_RET(weightContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 3456 | - inputTensor.weight = weightContiguous; | ||
| 3457 | - } | ||
| 3458 | 3454 | ||
| 3459 | // 检查conv3ddw确定性计算 | 3455 | // 检查conv3ddw确定性计算 |
| 3460 | if ((*outputMask)[1] && (input->GetViewShape().GetDimNum() == CONV3DINPUTDIM || | 3456 | if ((*outputMask)[1] && (input->GetViewShape().GetDimNum() == CONV3DINPUTDIM || |
| @@ -3622,27 +3618,23 @@ aclnnStatus aclnnConvTbcBackwardGetWorkspaceSize(const aclTensor* self, const ac | |||
| 3622 | Ops::NN::Conv::ConvTbcBackwardChecker convTbcBackwardChecker(inputTensor, outputTensor, tbcparams, npuArch); | 3618 | Ops::NN::Conv::ConvTbcBackwardChecker convTbcBackwardChecker(inputTensor, outputTensor, tbcparams, npuArch); |
| 3623 | auto ret = convTbcBackwardChecker.CheckTbcParams(); | 3619 | auto ret = convTbcBackwardChecker.CheckTbcParams(); |
| 3624 | CHECK_RET(ret == ACLNN_SUCCESS, ret); | 3620 | CHECK_RET(ret == ACLNN_SUCCESS, ret); |
| 3621 | + CHECK_RET(self != nullptr && input != nullptr && weight != nullptr && bias != nullptr, ACLNN_ERR_PARAM_NULLPTR); | ||
| 3625 | 3622 | ||
| 3626 | - if (self != nullptr) { | 3623 | + auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); |
| 3627 | - auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | 3624 | + CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3628 | - CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | 3625 | + inputTensor.self = selfContiguous; |
| 3629 | - inputTensor.self = selfContiguous; | 3626 | + |
| 3630 | - } | 3627 | + auto weightContiguous = l0op::Contiguous(weight, uniqueExecutor.get()); |
| 3631 | - if (weight != nullptr) { | 3628 | + CHECK_RET(weightContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3632 | - auto weightContiguous = l0op::Contiguous(weight, uniqueExecutor.get()); | 3629 | + inputTensor.weight = weightContiguous; |
| 3633 | - CHECK_RET(weightContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | 3630 | + |
| 3634 | - inputTensor.weight = weightContiguous; | 3631 | + auto inputContiguous = l0op::Contiguous(input, uniqueExecutor.get()); |
| 3635 | - } | 3632 | + CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3636 | - if (input != nullptr) { | 3633 | + inputTensor.input = inputContiguous; |
| 3637 | - auto inputContiguous = l0op::Contiguous(input, uniqueExecutor.get()); | 3634 | + |
| 3638 | - CHECK_RET(inputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | 3635 | + auto biasContiguous = l0op::Contiguous(bias, uniqueExecutor.get()); |
| 3639 | - inputTensor.input = inputContiguous; | 3636 | + CHECK_RET(biasContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 3640 | - } | 3637 | + inputTensor.bias = biasContiguous; |
| 3641 | - if (bias != nullptr) { | ||
| 3642 | - auto biasContiguous = l0op::Contiguous(bias, uniqueExecutor.get()); | ||
| 3643 | - CHECK_RET(biasContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 3644 | - inputTensor.bias = biasContiguous; | ||
| 3645 | - } | ||
| 3646 | 3638 | ||
| 3647 | // 设置param | 3639 | // 设置param |
| 3648 | FVector<int64_t> newStride = {1}; | 3640 | FVector<int64_t> newStride = {1}; |