已合并
修复random/drop_out_v3/op_api_dropout_v3.cpp代码风格问题 #3970
f2577359758创建于 7月10日
修复random/drop_out_v3/op_api_dropout_v3.cpp代码风格问题 #3970
已合并
f2577359758创建于 7月10日
1 个文件变更+16-17
@@ -17,7 +17,7 @@
17#include "aclnn_kernels/cast.h"17#include "aclnn_kernels/cast.h"
18#include "aclnn_kernels/contiguous.h"18#include "aclnn_kernels/contiguous.h"
19#include "dropout_v3.h"19#include "dropout_v3.h"
20-#include "conversion/fill//op_api/fill.h"20+#include "conversion/fill/op_api/fill.h"
21#include "math/zero_op/op_api/zero_op.h"21#include "math/zero_op/op_api/zero_op.h"
22#include "aclnn_kernels/common/op_error_check.h"22#include "aclnn_kernels/common/op_error_check.h"
23#include "op_api/aclnn_check.h"23#include "op_api/aclnn_check.h"
@@ -42,8 +42,8 @@ static const int64_t UINT8_BIT_NUMBER = 8;
42static const int8_t MAX_MASK_NUM = -1;42static const int8_t MAX_MASK_NUM = -1;
43 43 
44// 根据API定义,需要列出所能支持的所有dtype44// 根据API定义,需要列出所能支持的所有dtype
45-static const std::initializer_list<op::DataType> ASCEND910_DTYPE_SUPPORT_LIST = {45+static const std::initializer_list<op::DataType> ASCEND910_DTYPE_SUPPORT_LIST = {op::DataType::DT_FLOAT,
46- op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16};46+ op::DataType::DT_FLOAT16};
47 47 
48static const std::initializer_list<op::DataType> ARCH3510_DTYPE_SUPPORT_LIST = {48static const std::initializer_list<op::DataType> ARCH3510_DTYPE_SUPPORT_LIST = {
49 op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_BF16};49 op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_BF16};
@@ -122,9 +122,8 @@ static bool CheckShape(const aclTensor* input, const aclTensor* out, const aclTe
122 return true;122 return true;
123}123}
124 124 
125-static inline aclnnStatus CheckParams(125+static inline aclnnStatus CheckParams(const aclTensor* input, const aclTensor* optionalNoiseShape, double p,
126- const aclTensor* input, const aclTensor* optionalNoiseShape, double p, const aclTensor* out,126+ const aclTensor* out, const aclTensor* maskOut)
127- const aclTensor* maskOut)
128{127{
129 // 1. 检查参数是否为空指针128 // 1. 检查参数是否为空指针
130 CHECK_RET(CheckNotNull(input, out, maskOut), ACLNN_ERR_PARAM_NULLPTR);129 CHECK_RET(CheckNotNull(input, out, maskOut), ACLNN_ERR_PARAM_NULLPTR);
@@ -174,9 +173,9 @@ static inline const aclTensor* FillScalar(const aclTensor* out, int8_t val, aclO
174 return out;173 return out;
175}174}
176 175 
177-aclnnStatus aclnnDropoutV3GetWorkspaceSize(176+aclnnStatus aclnnDropoutV3GetWorkspaceSize(const aclTensor* input, const aclTensor* optionalNoiseShape, double p,
178- const aclTensor* input, const aclTensor* optionalNoiseShape, double p, int64_t seed, int64_t offset, aclTensor* out,177+ int64_t seed, int64_t offset, aclTensor* out, aclTensor* maskOut,
179- aclTensor* maskOut, uint64_t* workspaceSize, aclOpExecutor** executor)178+ uint64_t* workspaceSize, aclOpExecutor** executor)
180{179{
181 L2_DFX_PHASE_1(aclnnDropoutV3, DFX_IN(input, optionalNoiseShape, p, seed, offset), DFX_OUT(out, maskOut));180 L2_DFX_PHASE_1(aclnnDropoutV3, DFX_IN(input, optionalNoiseShape, p, seed, offset), DFX_OUT(out, maskOut));
182 // 固定写法,创建OpExecutor181 // 固定写法,创建OpExecutor
@@ -208,16 +207,16 @@ aclnnStatus aclnnDropoutV3GetWorkspaceSize(
208 maskResult = FillScalar(maskOut, 0, uniqueExecutor.get());207 maskResult = FillScalar(maskOut, 0, uniqueExecutor.get());
209 } else {208 } else {
210 FVector<double> probVector = {p};209 FVector<double> probVector = {p};
211- auto probTensor =210+ auto probTensor = uniqueExecutor.get()->ConvertToTensor(probVector.data(), probVector.size(),
212- uniqueExecutor.get()->ConvertToTensor(probVector.data(), probVector.size(), op::DataType::DT_DOUBLE);211+ op::DataType::DT_DOUBLE);
213 FVector<int64_t> seedVector = {seed};212 FVector<int64_t> seedVector = {seed};
214- auto seedTensor =213+ auto seedTensor = uniqueExecutor.get()->ConvertToTensor(seedVector.data(), seedVector.size(),
215- uniqueExecutor.get()->ConvertToTensor(seedVector.data(), seedVector.size(), op::DataType::DT_INT64);214+ op::DataType::DT_INT64);
216 FVector<int64_t> offsetVector = {0, offset};215 FVector<int64_t> offsetVector = {0, offset};
217- auto offsetTensor =216+ auto offsetTensor = uniqueExecutor.get()->ConvertToTensor(offsetVector.data(), offsetVector.size(),
218- uniqueExecutor.get()->ConvertToTensor(offsetVector.data(), offsetVector.size(), op::DataType::DT_INT64);217+ op::DataType::DT_INT64);
219- auto dropOutResult = l0op::DropoutV3(218+ auto dropOutResult = l0op::DropoutV3(inputContiguous, optionalNoiseShape, probTensor, seedTensor, offsetTensor,
220- inputContiguous, optionalNoiseShape, probTensor, seedTensor, offsetTensor, maskOut, uniqueExecutor.get());219+ maskOut, uniqueExecutor.get());
221 CHECK_RET(CheckTupleNullptr(dropOutResult), ACLNN_ERR_INNER_NULLPTR);220 CHECK_RET(CheckTupleNullptr(dropOutResult), ACLNN_ERR_INNER_NULLPTR);
222 outResult = std::get<0>(dropOutResult);221 outResult = std::get<0>(dropOutResult);
223 maskResult = std::get<1>(dropOutResult);222 maskResult = std::get<1>(dropOutResult);