已合并
delete extra level2_base_loss.h #5254
delete extra level2_base_loss.h #5254
已合并
li_wei21创建于 5月26日
共 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-#ifndef LEVEL2_BASE_LOSS_H_
12-#define LEVEL2_BASE_LOSS_H_
13- 
14-#include "aclnn_kernels/contiguous.h"
15-#include "aclnn/aclnn_base.h"
16-#include "opdev/shape_utils.h"
17-#include "level0/fill.h"
18-#include "aclnn_kernels/reshape.h"
19-#include "op_api/level2_base.h"
20- 
21-#ifdef __cplusplus
22-extern "C" {
23-#endif
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-#ifdef __cplusplus
86-}
87-#endif
88-#endif // LEVEL2_BASE_LOSS_H_
@@ -29,7 +29,7 @@
29#include "opdev/shape_utils.h"29#include "opdev/shape_utils.h"
30#include "opdev/tensor_view_utils.h"30#include "opdev/tensor_view_utils.h"
31#include "opdev/platform.h"31#include "opdev/platform.h"
32-#include "op_api/level2_base_loss.h"32+#include "loss/common/level2_base_loss.h"
33 33 
34using namespace op;34using namespace op;
35#ifdef __cplusplus35#ifdef __cplusplus
@@ -31,7 +31,7 @@
31#include "opdev/shape_utils.h"31#include "opdev/shape_utils.h"
32#include "opdev/tensor_view_utils.h"32#include "opdev/tensor_view_utils.h"
33#include "opdev/platform.h"33#include "opdev/platform.h"
34-#include "op_api/level2_base_loss.h"34+#include "loss/common/level2_base_loss.h"
35 35 
36using namespace op;36using namespace op;
37#ifdef __cplusplus37#ifdef __cplusplus
@@ -23,7 +23,7 @@
23#include "opdev/tensor_view_utils.h"23#include "opdev/tensor_view_utils.h"
24#include "aclnn_kernels/common/op_error_check.h"24#include "aclnn_kernels/common/op_error_check.h"
25#include "op_api/op_api_def.h"25#include "op_api/op_api_def.h"
26-#include "op_api/level2_base_loss.h"26+#include "loss/common/level2_base_loss.h"
27 27 
28using namespace op;28using namespace op;
29#ifdef __cplusplus29#ifdef __cplusplus