已合并
delete extra level2_base_loss.h #5254
li_wei21创建于 5月26日
delete extra level2_base_loss.h #5254
已合并
共 4 个文件变更+3-91
| @@ -1,88 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | - */ | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | -extern "C" { | ||
| 23 | - | ||
| 24 | - | ||
| 25 | -namespace op | ||
| 26 | -{ | ||
| 27 | -// 针对reduction!='none'且self为空tensor场景,对out按照给定值进行填充 | ||
| 28 | -inline static aclnnStatus CheckFillScalarLoss(aclTensor* out, float val, aclOpExecutor* executor) | ||
| 29 | -{ | ||
| 30 | - FVector<int64_t> tmp = {1}; | ||
| 31 | - auto dims = executor->ConvertToTensor(tmp.data(), tmp.size(), op::DataType::DT_INT64); | ||
| 32 | - auto shapeArray = executor->AllocIntArray(tmp.data(), tmp.size()); | ||
| 33 | - CHECK_RET(shapeArray != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 34 | - | ||
| 35 | - FVector<float> valVector = {val}; | ||
| 36 | - auto valTensor = executor->ConvertToTensor(valVector.data(), valVector.size(), out->GetDataType()); | ||
| 37 | - auto fillOut = l0op::Fill(dims, valTensor, shapeArray, executor); | ||
| 38 | - CHECK_RET(fillOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 39 | - auto viewCopyResult = l0op::ViewCopy(fillOut, out, executor); | ||
| 40 | - CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 41 | - return ACLNN_SUCCESS; | ||
| 42 | -} | ||
| 43 | - | ||
| 44 | -inline static aclIntArray* GetBroadcastShapeLossBackward(const op::Shape broadcastShape, aclOpExecutor* executor) | ||
| 45 | -{ | ||
| 46 | - int64_t tensorSize = static_cast<int64_t>(broadcastShape.GetDimNum()); | ||
| 47 | - std::vector<int64_t> tensorShape(tensorSize); | ||
| 48 | - for (int i = 0; i < tensorSize; i++) { | ||
| 49 | - tensorShape[i] = broadcastShape[i]; | ||
| 50 | - } | ||
| 51 | - return executor->AllocIntArray(tensorShape.data(), tensorSize); | ||
| 52 | -} | ||
| 53 | - | ||
| 54 | -inline static bool CheckDtypeValidMseLoss(const aclTensor* self, const aclTensor* target, const aclTensor* out, | ||
| 55 | - const std::initializer_list<op::DataType>& l1, | ||
| 56 | - const std::initializer_list<op::DataType>& l2) | ||
| 57 | -{ | ||
| 58 | - auto supportList = GetDtypeSupportListV2(l1, l2); | ||
| 59 | - // 检查self和target做数据类型推导后的数据类型是否在支持列表内 | ||
| 60 | - op::DataType promoteType = op::PromoteType(self->GetDataType(), target->GetDataType()); | ||
| 61 | - if (!CheckType(promoteType, supportList)) { | ||
| 62 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, | ||
| 63 | - "Expected self dtype [%s] and target dtype [%s] to be promotable but check failed.", | ||
| 64 | - ToString(self->GetDataType()).GetString(), ToString(target->GetDataType()).GetString()); | ||
| 65 | - return false; | ||
| 66 | - } | ||
| 67 | - | ||
| 68 | - // 检查self, target和out的数据类型是否在支持列表内 | ||
| 69 | - OP_CHECK_DTYPE_NOT_SUPPORT(self, supportList, return false); | ||
| 70 | - OP_CHECK_DTYPE_NOT_SUPPORT(target, supportList, return false); | ||
| 71 | - OP_CHECK_DTYPE_NOT_SUPPORT(out, supportList, return false); | ||
| 72 | - return true; | ||
| 73 | -} | ||
| 74 | - | ||
| 75 | -inline static bool CheckReductionMseLoss(int64_t reduction, int64_t REDUCTION_SUM_NUM, int64_t REDUCTION_NONE_NUM) | ||
| 76 | -{ | ||
| 77 | - if (reduction > REDUCTION_SUM_NUM || reduction < REDUCTION_NONE_NUM) { | ||
| 78 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Expected reduction to be between 0 and 2, but got %ld.", reduction); | ||
| 79 | - return false; | ||
| 80 | - } | ||
| 81 | - return true; | ||
| 82 | -} | ||
| 83 | - | ||
| 84 | -} // namespace op | ||
| 85 | - | ||
| 86 | -} | ||
| 87 | - | ||
| 88 | - | ||
| @@ -29,7 +29,7 @@ | |||
| 29 | 29 | ||
| 30 | 30 | ||
| 31 | 31 | ||
| 32 | -#include "op_api/level2_base_loss.h" | 32 | +#include "loss/common/level2_base_loss.h" |
| 33 | 33 | ||
| 34 | using namespace op; | 34 | using namespace op; |
| 35 | 35 | ||
| @@ -31,7 +31,7 @@ | |||
| 31 | 31 | ||
| 32 | 32 | ||
| 33 | 33 | ||
| 34 | -#include "op_api/level2_base_loss.h" | 34 | +#include "loss/common/level2_base_loss.h" |
| 35 | 35 | ||
| 36 | using namespace op; | 36 | using namespace op; |
| 37 | 37 | ||
| @@ -23,7 +23,7 @@ | |||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | -#include "op_api/level2_base_loss.h" | 26 | +#include "loss/common/level2_base_loss.h" |
| 27 | 27 | ||
| 28 | using namespace op; | 28 | using namespace op; |
| 29 | 29 | ||