diff --git a/mindspore/ccsrc/backend/kernel_compiler/cpu/nnacl/base/cast_base.h b/mindspore/ccsrc/backend/kernel_compiler/cpu/nnacl/base/cast_base.h index 81044f6c54..5e31d8377c 100644 --- a/mindspore/ccsrc/backend/kernel_compiler/cpu/nnacl/base/cast_base.h +++ b/mindspore/ccsrc/backend/kernel_compiler/cpu/nnacl/base/cast_base.h @@ -13,8 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -#ifndef MINDSPORE_NNACL_CAST_BASE_H_ -#define MINDSPORE_NNACL_CAST_BASE_H_ +#ifndef MINDSPORE_CCSRC_BACKEND_KERNEL_COMPILER_CPU_NNACL_BASE_CAST_BASE_H_ +#define MINDSPORE_CCSRC_BACKEND_KERNEL_COMPILER_CPU_NNACL_BASE_CAST_BASE_H_ #include "nnacl/op_base.h" #include "nnacl/nnacl_common.h" @@ -71,8 +71,19 @@ inline void Uint8ToFp16(const uint8_t *input, float16_t *output, int number) { output[i] = (float16_t)input[i]; } } -#endif +inline void Float32ToFp16(const float *input, float16_t *output, int number) { + for (int i = 0; i < number; ++i) { + output[i] = (float16_t)(input[i]); + } +} + +inline void Fp16ToFloat32(const float16_t *input, float *output, int number) { + for (int i = 0; i < number; ++i) { + output[i] = (float)(input[i]); + } +} +#else inline void Fp16ToFloat32(const uint16_t *input, float *output, int number) { for (int i = 0; i < number; ++i) { output[i] = ShortToFloat32(input[i]); @@ -84,6 +95,7 @@ inline void Float32ToFp16(const float *input, uint16_t *output, int number) { output[i] = Float32ToShort(input[i]); } } +#endif inline void Float32ToInt32(const float *input, int32_t *output, int number) { for (int i = 0; i < number; ++i) { @@ -125,4 +137,4 @@ inline void Float32ToBool(const float *input, bool *output, int number) { } #endif -#endif // MINDSPORE_NNACL_CAST_BASE_H_ +#endif // MINDSPORE_CCSRC_BACKEND_KERNEL_COMPILER_CPU_NNACL_BASE_CAST_BASE_H_ diff --git a/mindspore/lite/examples/train_lenet/src/net_runner.cc b/mindspore/lite/examples/train_lenet/src/net_runner.cc index 29e9d3fb42..0180eca5ae 100644 --- a/mindspore/lite/examples/train_lenet/src/net_runner.cc +++ b/mindspore/lite/examples/train_lenet/src/net_runner.cc @@ -151,7 +151,6 @@ void NetRunner::InitAndFigureInputs() { context.thread_num_ = 2; mindspore::lite::TrainCfg train_cfg; - train_cfg.mix_precision_cfg_.is_raw_mix_precision_ = is_raw_mix_precision_; session_ = mindspore::session::TrainSession::CreateTrainSession(ms_file_, &context, true, &train_cfg); MS_ASSERT(session_ != nullptr); @@ -257,7 +256,7 @@ void NetRunner::Usage() { bool NetRunner::ReadArgs(int argc, char *argv[]) { int opt; - while ((opt = getopt(argc, argv, "f:e:d:s:ihc:vmob:")) != -1) { + while ((opt = getopt(argc, argv, "f:e:d:s:ihc:vob:")) != -1) { switch (opt) { case 'f': ms_file_ = std::string(optarg); @@ -280,9 +279,6 @@ bool NetRunner::ReadArgs(int argc, char *argv[]) { case 'b': virtual_batch_ = atoi(optarg); break; - case 'm': - is_raw_mix_precision_ = true; - break; case 'h': default: Usage(); diff --git a/mindspore/lite/examples/train_lenet/src/net_runner.h b/mindspore/lite/examples/train_lenet/src/net_runner.h index 433dbfa700..b02f3620b2 100644 --- a/mindspore/lite/examples/train_lenet/src/net_runner.h +++ b/mindspore/lite/examples/train_lenet/src/net_runner.h @@ -63,7 +63,6 @@ class NetRunner { int batch_size_ = 32; int h_ = 32; int w_ = 32; - bool is_raw_mix_precision_ = false; }; #endif // MINDSPORE_LITE_EXAMPLES_TRAIN_LENET_SRC_NET_RUNNER_H_ diff --git a/mindspore/lite/include/train/train_cfg.h b/mindspore/lite/include/train/train_cfg.h index 5c0246b965..853da49733 100644 --- a/mindspore/lite/include/train/train_cfg.h +++ b/mindspore/lite/include/train/train_cfg.h @@ -29,28 +29,24 @@ class MixPrecisionCfg { this->loss_scale_ = 128.0f; this->keep_batchnorm_fp32_ = true; this->num_of_not_nan_iter_th_ = 1000; - this->is_raw_mix_precision_ = false; } MixPrecisionCfg(const MixPrecisionCfg &rhs) { this->dynamic_loss_scale_ = rhs.dynamic_loss_scale_; this->loss_scale_ = rhs.loss_scale_; this->keep_batchnorm_fp32_ = rhs.keep_batchnorm_fp32_; this->num_of_not_nan_iter_th_ = rhs.num_of_not_nan_iter_th_; - this->is_raw_mix_precision_ = rhs.is_raw_mix_precision_; } MixPrecisionCfg &operator=(MixPrecisionCfg const &rhs) { this->dynamic_loss_scale_ = rhs.dynamic_loss_scale_; this->loss_scale_ = rhs.loss_scale_; this->keep_batchnorm_fp32_ = rhs.keep_batchnorm_fp32_; this->num_of_not_nan_iter_th_ = rhs.num_of_not_nan_iter_th_; - this->is_raw_mix_precision_ = rhs.is_raw_mix_precision_; return *this; } - bool dynamic_loss_scale_ = false; /**< Enable\disable dynamic loss scale during mix precision training */ - float loss_scale_; /**< Initial loss scale factor */ - bool keep_batchnorm_fp32_ = true; /**< Keep batch norm in FP32 while training */ - uint32_t num_of_not_nan_iter_th_; /**< a threshold for modifying loss scale when dynamic loss scale is enabled */ - bool is_raw_mix_precision_ = false; /**< Is mix precision model export from mindspore */ + bool dynamic_loss_scale_ = false; /**< Enable\disable dynamic loss scale during mix precision training */ + float loss_scale_; /**< Initial loss scale factor */ + bool keep_batchnorm_fp32_ = true; /**< Keep batch norm in FP32 while training */ + uint32_t num_of_not_nan_iter_th_; /**< a threshold for modifying loss scale when dynamic loss scale is enabled */ }; /// \brief TrainCfg defined for holding train configuration. diff --git a/mindspore/lite/src/cxx_api/train/converters.cc b/mindspore/lite/src/cxx_api/train/converters.cc index 5c63fa8fd0..46b4c3064c 100644 --- a/mindspore/lite/src/cxx_api/train/converters.cc +++ b/mindspore/lite/src/cxx_api/train/converters.cc @@ -36,7 +36,6 @@ Status A2L_ConvertConfig(const TrainCfg *a_train_cfg, lite::TrainCfg *l_train_cf l_train_cfg->mix_precision_cfg_.loss_scale_ = a_train_cfg->mix_precision_cfg_.loss_scale_; l_train_cfg->mix_precision_cfg_.keep_batchnorm_fp32_ = (a_train_cfg->optimization_level_ != kO3); l_train_cfg->mix_precision_cfg_.num_of_not_nan_iter_th_ = a_train_cfg->mix_precision_cfg_.num_of_not_nan_iter_th_; - l_train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = a_train_cfg->mix_precision_cfg_.is_raw_mix_precision_; l_train_cfg->accumulate_gradients_ = a_train_cfg->accumulate_gradients_; return kSuccess; } diff --git a/mindspore/lite/src/mindrt_executor.cc b/mindspore/lite/src/mindrt_executor.cc index 0953af9413..40f0438c44 100644 --- a/mindspore/lite/src/mindrt_executor.cc +++ b/mindspore/lite/src/mindrt_executor.cc @@ -136,17 +136,20 @@ void MindrtExecutor::TransferGraphOutput() { auto src_tensor = tensor_map.first; dst_tensor->set_shape(src_tensor->shape()); /* dst tensor free in FreeOutputTensor */ - +#ifdef ENABLE_FP16 if (src_tensor->data_type() == kNumberTypeFloat16) { dst_tensor->MallocData(); - Fp16ToFloat32(reinterpret_cast(src_tensor->MutableData()), + Fp16ToFloat32(reinterpret_cast(src_tensor->MutableData()), reinterpret_cast(dst_tensor->data()), dst_tensor->ElementsNum()); } else { +#endif dst_tensor->set_data(src_tensor->data()); if (IS_RUNTIME_ALLOCATOR(src_tensor->allocator()) == false) { src_tensor->set_data(nullptr); } +#ifdef ENABLE_FP16 } +#endif src_tensor->DecRefCount(); } return; diff --git a/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_1x1_fp16.cc b/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_1x1_fp16.cc index 0d5b596268..3bbf6824a4 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_1x1_fp16.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_1x1_fp16.cc @@ -157,10 +157,12 @@ int Convolution1x1FP16CPUKernel::Prepare() { size_t size = input_channel * UP_ROUND(output_channel, col_tile_) * sizeof(float16_t); set_workspace_size(size); } - matmul_param_ = new (std::nothrow) MatMulParameter(); if (matmul_param_ == nullptr) { - MS_LOG(ERROR) << "Init matmul_param_ failed."; - return RET_ERROR; + matmul_param_ = new (std::nothrow) MatMulParameter(); + if (matmul_param_ == nullptr) { + MS_LOG(ERROR) << "Init matmul_param_ failed."; + return RET_ERROR; + } } int ret = InitConvWeightBias(); if (ret != RET_OK) { diff --git a/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_winograd_fp16.cc b/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_winograd_fp16.cc index b0d4b14a77..eea2a3691c 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_winograd_fp16.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp16/convolution_winograd_fp16.cc @@ -150,6 +150,10 @@ int ConvolutionWinogradFP16CPUKernel::Prepare() { #else row_tile_ = C12NUM; #endif + kernel_unit_ = conv_param_->kernel_h_; + input_unit_ = output_unit_ + kernel_unit_ - 1; + conv_param_->input_unit_ = input_unit_; + conv_param_->output_unit_ = output_unit_; if (op_parameter_->is_train_session_) { auto weight_tensor = in_tensors_.at(kWeightIndex); CHECK_NULL_RETURN(weight_tensor); @@ -159,10 +163,6 @@ int ConvolutionWinogradFP16CPUKernel::Prepare() { auto trans_matrix_data_size = input_unit_ * input_unit_ * in_channel * oc_block_num * col_tile_ * sizeof(float16_t); set_workspace_size(trans_matrix_data_size); } - kernel_unit_ = conv_param_->kernel_h_; - input_unit_ = output_unit_ + kernel_unit_ - 1; - conv_param_->input_unit_ = input_unit_; - conv_param_->output_unit_ = output_unit_; auto ret = InitConvWeightBias(); if (ret != RET_OK) { diff --git a/mindspore/lite/src/runtime/kernel/arm/fp16/deconvolution_fp16.cc b/mindspore/lite/src/runtime/kernel/arm/fp16/deconvolution_fp16.cc index 324a4d9e4f..a5f3d218e4 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp16/deconvolution_fp16.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp16/deconvolution_fp16.cc @@ -199,10 +199,12 @@ int DeConvolutionFp16CPUKernel::Prepare() { size_t weight_pack_size = input_channel * kernel_w * kernel_h * UP_ROUND(output_channel, C8NUM) * sizeof(float16_t); set_workspace_size(weight_pack_size); } - matmul_param_ = new (std::nothrow) MatMulParameter(); if (matmul_param_ == nullptr) { - MS_LOG(ERROR) << "Memory allocation failed"; - return RET_ERROR; + matmul_param_ = new (std::nothrow) MatMulParameter(); + if (matmul_param_ == nullptr) { + MS_LOG(ERROR) << "Memory allocation failed"; + return RET_ERROR; + } } int ret = InitConvWeightBias(); if (ret != RET_OK) { diff --git a/mindspore/lite/src/runtime/kernel/arm/fp32/cast_fp32.cc b/mindspore/lite/src/runtime/kernel/arm/fp32/cast_fp32.cc index 6ee360eb69..97c2071b08 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp32/cast_fp32.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp32/cast_fp32.cc @@ -75,9 +75,14 @@ int CastCPUKernel::CastToFp32(const lite::Tensor *input, lite::Tensor *output, i reinterpret_cast(output_data) + offset, data_num); break; case kNumberTypeFloat16: +#ifdef ENABLE_FP16 + Fp16ToFloat32(reinterpret_cast(input->data()) + offset, + reinterpret_cast(output_data) + offset, data_num); +#else Fp16ToFloat32(reinterpret_cast(input->data()) + offset, reinterpret_cast(output_data) + offset, data_num); break; +#endif case kNumberTypeInt64: Int64ToFloat32(reinterpret_cast(input->data()) + offset, reinterpret_cast(output_data) + offset, data_num); @@ -93,11 +98,11 @@ int CastCPUKernel::CastToFp16(const lite::Tensor *input, lite::Tensor *output, i auto input_data_type = input->data_type(); auto output_data = output->data(); switch (input_data_type) { +#ifdef ENABLE_FP16 case kNumberTypeFloat32: Float32ToFp16(reinterpret_cast(input->data()) + offset, - reinterpret_cast(output_data) + offset, data_num); + reinterpret_cast(output_data) + offset, data_num); break; -#ifdef ENABLE_FP16 case kNumberTypeInt64: Int64ToFp16(reinterpret_cast(input->data()) + offset, reinterpret_cast(output_data) + offset, data_num); @@ -113,6 +118,11 @@ int CastCPUKernel::CastToFp16(const lite::Tensor *input, lite::Tensor *output, i Uint8ToFp16(reinterpret_cast(input->data()) + offset, reinterpret_cast(output_data) + offset, data_num); break; +#else + case kNumberTypeFloat32: + Float32ToFp16(reinterpret_cast(input->data()) + offset, + reinterpret_cast(output_data) + offset, data_num); + break; #endif default: MS_LOG(ERROR) << "Unsupported input data type " << input_data_type; diff --git a/mindspore/lite/src/runtime/kernel/arm/fp32/convolution_1x1_fp32.cc b/mindspore/lite/src/runtime/kernel/arm/fp32/convolution_1x1_fp32.cc index c97b10e7b1..3b35601ab3 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp32/convolution_1x1_fp32.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp32/convolution_1x1_fp32.cc @@ -111,10 +111,12 @@ int Convolution1x1CPUKernel::Prepare() { row_tile_ = C12NUM; col_tile_ = C8NUM; #endif - matmul_param_ = new (std::nothrow) MatMulParameter; if (matmul_param_ == nullptr) { - MS_LOG(ERROR) << "Memory allocation failed"; - return RET_ERROR; + matmul_param_ = new (std::nothrow) MatMulParameter; + if (matmul_param_ == nullptr) { + MS_LOG(ERROR) << "Memory allocation failed"; + return RET_ERROR; + } } if (op_parameter_->is_train_session_) { auto filter_tensor = in_tensors_.at(kWeightIndex); diff --git a/mindspore/lite/src/runtime/kernel/arm/fp32/deconvolution_fp32.cc b/mindspore/lite/src/runtime/kernel/arm/fp32/deconvolution_fp32.cc index 32e30b9adc..6fc10ac527 100644 --- a/mindspore/lite/src/runtime/kernel/arm/fp32/deconvolution_fp32.cc +++ b/mindspore/lite/src/runtime/kernel/arm/fp32/deconvolution_fp32.cc @@ -183,10 +183,12 @@ int DeConvolutionCPUKernel::Prepare() { size_t pack_weight_size = input_channel * kernel_w_ * kernel_h_ * output_aligned_size * sizeof(float); set_workspace_size(pack_weight_size); } - matmul_param_ = new (std::nothrow) MatMulParameter(); if (matmul_param_ == nullptr) { - MS_LOG(ERROR) << "Memory allocation failed"; - return RET_ERROR; + matmul_param_ = new (std::nothrow) MatMulParameter(); + if (matmul_param_ == nullptr) { + MS_LOG(ERROR) << "Memory allocation failed"; + return RET_ERROR; + } } if (in_tensors_.at(kWeightIndex)->data() != nullptr) { int error_code = InitConvWeightBias(); diff --git a/mindspore/lite/src/scheduler.cc b/mindspore/lite/src/scheduler.cc index 7398922cb2..a87c42b0af 100644 --- a/mindspore/lite/src/scheduler.cc +++ b/mindspore/lite/src/scheduler.cc @@ -1637,7 +1637,7 @@ void Scheduler::SetKernelTensorDataType(kernel::LiteKernel *kernel) { } } for (auto tensor : kernel->out_tensors()) { - if (tensor->data_type() == kNumberTypeFloat16) { + if (tensor->data_type() == kNumberTypeFloat16 && kernel->type() != schema::PrimitiveType_Cast) { tensor->set_data_type(kNumberTypeFloat32); } } diff --git a/mindspore/lite/src/train/train_populate_parameter.cc b/mindspore/lite/src/train/train_populate_parameter.cc index f3164580b9..9b48f72ba5 100644 --- a/mindspore/lite/src/train/train_populate_parameter.cc +++ b/mindspore/lite/src/train/train_populate_parameter.cc @@ -249,17 +249,6 @@ OpParameter *PopulateAvgPoolGradParameter(const void *prim) { } pooling_param->round_mode_ = RoundMode_No; pooling_param->pool_mode_ = PoolMode_AvgPool; - switch (value->pad_mode()) { - case schema::PadMode_SAME: - pooling_param->pad_mode_ = Pad_same; - break; - case schema::PadMode_VALID: - pooling_param->pad_mode_ = Pad_valid; - break; - default: - pooling_param->pad_mode_ = Pad_pad; - break; - } return reinterpret_cast(pooling_param); } @@ -293,8 +282,8 @@ OpParameter *PopulateConvolutionGradFilterParameter(const void *prim) { param->kernel_h_ = value->kernel_size()->Get(0); param->kernel_w_ = value->kernel_size()->Get(1); - param->stride_h_ = value->stride()->Get(0); - param->stride_w_ = value->stride()->Get(1); + param->stride_h_ = value->stride()->Get((value->stride()->size()) - 2); + param->stride_w_ = value->stride()->Get((value->stride()->size()) - 1); param->dilation_h_ = value->dilation()->Get(0); param->dilation_w_ = value->dilation()->Get(1); param->pad_u_ = value->pad_list()->Get(0); @@ -313,6 +302,17 @@ OpParameter *PopulateConvolutionGradFilterParameter(const void *prim) { default: break; } + switch (value->pad_mode()) { + case schema::PadMode_SAME: + param->pad_mode_ = Pad_same; + break; + case schema::PadMode_VALID: + param->pad_mode_ = Pad_valid; + break; + default: + param->pad_mode_ = Pad_pad; + break; + } return reinterpret_cast(param); } @@ -331,8 +331,8 @@ OpParameter *PopulateConvolutionGradInputParameter(const void *prim) { param->kernel_h_ = value->kernel_size()->Get(0); param->kernel_w_ = value->kernel_size()->Get(1); - param->stride_h_ = value->stride()->Get(0); - param->stride_w_ = value->stride()->Get(1); + param->stride_h_ = value->stride()->Get((value->stride()->size()) - 2); + param->stride_w_ = value->stride()->Get((value->stride()->size()) - 1); param->dilation_h_ = value->dilation()->Get(0); param->dilation_w_ = value->dilation()->Get(1); param->pad_u_ = value->pad_list()->Get(0); @@ -351,6 +351,17 @@ OpParameter *PopulateConvolutionGradInputParameter(const void *prim) { default: break; } + switch (value->pad_mode()) { + case schema::PadMode_SAME: + param->pad_mode_ = Pad_same; + break; + case schema::PadMode_VALID: + param->pad_mode_ = Pad_valid; + break; + default: + param->pad_mode_ = Pad_pad; + break; + } return reinterpret_cast(param); } @@ -514,68 +525,57 @@ OpParameter *PopulateLstmGradParameter(const void *prim) { } void PopulateTrainParameters() { - lite::Registry ApplyMomentumParameterRegistry(schema::PrimitiveType_ApplyMomentum, PopulateApplyMomentumParameter, - lite::SCHEMA_CUR); - lite::Registry BiasGradParameterRegistry(schema::PrimitiveType_BiasAddGrad, PopulateBiasGradParameter, - lite::SCHEMA_CUR); - lite::Registry SoftmaxCrossEntropyParameterRegistry(schema::PrimitiveType_SoftmaxCrossEntropyWithLogits, - PopulateSoftmaxCrossEntropyParameter, lite::SCHEMA_CUR); - lite::Registry SparseSoftmaxCrossEntropyParameterRegistry(schema::PrimitiveType_SparseSoftmaxCrossEntropyWithLogits, - PopulateSparseSoftmaxCrossEntropyWithLogitsParameter, - lite::SCHEMA_CUR); - lite::Registry ActivationParameterRegistry(schema::PrimitiveType_ActivationGrad, PopulateActivationGradParameter, - lite::SCHEMA_CUR); - lite::Registry DependParameterRegistry(schema::PrimitiveType_Depend, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry Conv2DGradFilterParameterRegistry(schema::PrimitiveType_Conv2DBackpropFilterFusion, - PopulateConvolutionGradFilterParameter, lite::SCHEMA_CUR); - lite::Registry Conv2DGradInputParameterRegistry(schema::PrimitiveType_Conv2DBackpropInputFusion, - PopulateConvolutionGradInputParameter, lite::SCHEMA_CUR); - lite::Registry avgPoolParameterRegistry(schema::PrimitiveType_AvgPoolGrad, PopulateAvgPoolGradParameter, + Registry ApplyMomentumParameterRegistry(schema::PrimitiveType_ApplyMomentum, PopulateApplyMomentumParameter, lite::SCHEMA_CUR); - lite::Registry maxPoolParameterRegistry(schema::PrimitiveType_MaxPoolGrad, PopulateMaxPoolGradParameter, - lite::SCHEMA_CUR); - lite::Registry PowerGradParameterRegistry(schema::PrimitiveType_PowerGrad, PopulatePowerGradParameter, - lite::SCHEMA_CUR); - lite::Registry SgdParameterRegistry(schema::PrimitiveType_SGD, PopulateSgdParameter, lite::SCHEMA_CUR); - lite::Registry BNGradParameterRegistry(schema::PrimitiveType_BatchNormGrad, PopulateBNGradParameter, - lite::SCHEMA_CUR); - lite::Registry AdamParameterRegistry(schema::PrimitiveType_Adam, PopulateAdamParameter, lite::SCHEMA_CUR); - lite::Registry AssignParameterRegistry(schema::PrimitiveType_Assign, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry AssignAddParameterRegistry(schema::PrimitiveType_AssignAdd, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry BinaryCrossEntropyParameterRegistry(schema::PrimitiveType_BinaryCrossEntropy, PopulateBCEParameter, - lite::SCHEMA_CUR); - lite::Registry BinaryCrossEntropyGradParameterRegistry(schema::PrimitiveType_BinaryCrossEntropyGrad, - PopulateBCEGradParameter, lite::SCHEMA_CUR); - lite::Registry OnesLikeParameterRegistry(schema::PrimitiveType_OnesLike, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry UnsortedSegmentSumParameterRegistry(schema::PrimitiveType_UnsortedSegmentSum, - lite::DefaultPopulateParameter, lite::SCHEMA_CUR); - lite::Registry DropoutParameterRegistry(schema::PrimitiveType_Dropout, PopulateDropoutParameter, lite::SCHEMA_CUR); - lite::Registry DropGradParameterRegistry(schema::PrimitiveType_DropoutGrad, PopulateDropoutGradParameter, - lite::SCHEMA_CUR); - lite::Registry MaximumGradParameterRegistry(schema::PrimitiveType_MaximumGrad, PopulateArithmeticGradParameter, - lite::SCHEMA_CUR); - lite::Registry MinimumGradParameterRegistry(schema::PrimitiveType_MinimumGrad, PopulateArithmeticGradParameter, - lite::SCHEMA_CUR); - lite::Registry SmoothL1LossRegistry(schema::PrimitiveType_SmoothL1Loss, PopulateSmoothL1LossParameter, + Registry BiasGradParameterRegistry(schema::PrimitiveType_BiasAddGrad, PopulateBiasGradParameter, lite::SCHEMA_CUR); + Registry SoftmaxCrossEntropyParameterRegistry(schema::PrimitiveType_SoftmaxCrossEntropyWithLogits, + PopulateSoftmaxCrossEntropyParameter, lite::SCHEMA_CUR); + Registry SparseSoftmaxCrossEntropyParameterRegistry(schema::PrimitiveType_SparseSoftmaxCrossEntropyWithLogits, + PopulateSparseSoftmaxCrossEntropyWithLogitsParameter, + lite::SCHEMA_CUR); + Registry ActivationParameterRegistry(schema::PrimitiveType_ActivationGrad, PopulateActivationGradParameter, + lite::SCHEMA_CUR); + Registry DependParameterRegistry(schema::PrimitiveType_Depend, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry Conv2DGradFilterParameterRegistry(schema::PrimitiveType_Conv2DBackpropFilterFusion, + PopulateConvolutionGradFilterParameter, lite::SCHEMA_CUR); + Registry Conv2DGradInputParameterRegistry(schema::PrimitiveType_Conv2DBackpropInputFusion, + PopulateConvolutionGradInputParameter, lite::SCHEMA_CUR); + Registry avgPoolParameterRegistry(schema::PrimitiveType_AvgPoolGrad, PopulateAvgPoolGradParameter, lite::SCHEMA_CUR); + Registry maxPoolParameterRegistry(schema::PrimitiveType_MaxPoolGrad, PopulateMaxPoolGradParameter, lite::SCHEMA_CUR); + Registry PowerGradParameterRegistry(schema::PrimitiveType_PowerGrad, PopulatePowerGradParameter, lite::SCHEMA_CUR); + Registry SgdParameterRegistry(schema::PrimitiveType_SGD, PopulateSgdParameter, lite::SCHEMA_CUR); + Registry BNGradParameterRegistry(schema::PrimitiveType_BatchNormGrad, PopulateBNGradParameter, lite::SCHEMA_CUR); + Registry AdamParameterRegistry(schema::PrimitiveType_Adam, PopulateAdamParameter, lite::SCHEMA_CUR); + Registry AssignParameterRegistry(schema::PrimitiveType_Assign, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry AssignAddParameterRegistry(schema::PrimitiveType_AssignAdd, lite::DefaultPopulateParameter, + lite::SCHEMA_CUR); + Registry BinaryCrossEntropyParameterRegistry(schema::PrimitiveType_BinaryCrossEntropy, PopulateBCEParameter, + lite::SCHEMA_CUR); + Registry BinaryCrossEntropyGradParameterRegistry(schema::PrimitiveType_BinaryCrossEntropyGrad, + PopulateBCEGradParameter, lite::SCHEMA_CUR); + Registry OnesLikeParameterRegistry(schema::PrimitiveType_OnesLike, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry UnsortedSegmentSumParameterRegistry(schema::PrimitiveType_UnsortedSegmentSum, lite::DefaultPopulateParameter, + lite::SCHEMA_CUR); + Registry DropoutParameterRegistry(schema::PrimitiveType_Dropout, PopulateDropoutParameter, lite::SCHEMA_CUR); + Registry DropGradParameterRegistry(schema::PrimitiveType_DropoutGrad, PopulateDropoutGradParameter, lite::SCHEMA_CUR); + Registry MaximumGradParameterRegistry(schema::PrimitiveType_MaximumGrad, PopulateArithmeticGradParameter, + lite::SCHEMA_CUR); + Registry MinimumGradParameterRegistry(schema::PrimitiveType_MinimumGrad, PopulateArithmeticGradParameter, + lite::SCHEMA_CUR); + Registry SmoothL1LossRegistry(schema::PrimitiveType_SmoothL1Loss, PopulateSmoothL1LossParameter, lite::SCHEMA_CUR); + Registry SmoothL1LossGradRegistry(schema::PrimitiveType_SmoothL1LossGrad, PopulateSmoothL1LossGradParameter, + lite::SCHEMA_CUR); + Registry SigmoidCrossEntropyWithLogitsRegistry(schema::PrimitiveType_SigmoidCrossEntropyWithLogits, + lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry SigmoidCrossEntropyWithLogitsGradRegistry(schema::PrimitiveType_SigmoidCrossEntropyWithLogitsGrad, + lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry FlattenGradParameterRegistry(schema::PrimitiveType_FlattenGrad, lite::DefaultPopulateParameter, + lite::SCHEMA_CUR); + Registry StridedSliceGradParameterRegistry(schema::PrimitiveType_StridedSliceGrad, PopulateStridedSliceGradParameter, + lite::SCHEMA_CUR); + Registry SqrtGradParameterRegistry(schema::PrimitiveType_SqrtGrad, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); + Registry RsqrtGradParameterRegistry(schema::PrimitiveType_RsqrtGrad, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); - lite::Registry SmoothL1LossGradRegistry(schema::PrimitiveType_SmoothL1LossGrad, PopulateSmoothL1LossGradParameter, - lite::SCHEMA_CUR); - lite::Registry SigmoidCrossEntropyWithLogitsRegistry(schema::PrimitiveType_SigmoidCrossEntropyWithLogits, - lite::DefaultPopulateParameter, lite::SCHEMA_CUR); - lite::Registry SigmoidCrossEntropyWithLogitsGradRegistry(schema::PrimitiveType_SigmoidCrossEntropyWithLogitsGrad, - lite::DefaultPopulateParameter, lite::SCHEMA_CUR); - lite::Registry FlattenGradParameterRegistry(schema::PrimitiveType_FlattenGrad, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry StridedSliceGradParameterRegistry(schema::PrimitiveType_StridedSliceGrad, - PopulateStridedSliceGradParameter, lite::SCHEMA_CUR); - lite::Registry SqrtGradParameterRegistry(schema::PrimitiveType_SqrtGrad, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); - lite::Registry RsqrtGradParameterRegistry(schema::PrimitiveType_RsqrtGrad, lite::DefaultPopulateParameter, - lite::SCHEMA_CUR); Registry ResizeGradParameterRegistry(schema::PrimitiveType_ResizeGrad, PopulateResizeGradParameter, lite::SCHEMA_CUR); Registry AbsGradParameterRegistry(schema::PrimitiveType_AbsGrad, lite::DefaultPopulateParameter, lite::SCHEMA_CUR); Registry LSTMGradParameterRegistry(schema::PrimitiveType_LSTMGrad, PopulateLstmGradParameter, lite::SCHEMA_CUR); diff --git a/mindspore/lite/src/train/train_session.cc b/mindspore/lite/src/train/train_session.cc index ea524b4ff2..16af68be29 100644 --- a/mindspore/lite/src/train/train_session.cc +++ b/mindspore/lite/src/train/train_session.cc @@ -125,29 +125,26 @@ int TrainSession::InitCallBack() { if (!context_->IsCpuFloat16Enabled()) { return false; } - if (cfg_.mix_precision_cfg_.is_raw_mix_precision_) { - auto out_tensor_indexs = node->output_indices_; - if (out_tensor_indexs.empty()) { - MS_LOG(DEBUG) << "Debug: " << node->name_ << " fp32"; - return false; - } - auto is_fp16 = model_->all_tensors_.at(out_tensor_indexs[0])->dataType() == kNumberTypeFloat16; - MS_LOG(DEBUG) << "Debug: " << node->name_ << ((is_fp16) ? " fp16" : " fp32"); - return is_fp16; - } + bool force_fp16 = false; auto node_type = GetPrimitiveType(node->primitive_, SCHEMA_VERSION::SCHEMA_CUR); if (node_type == schema::PrimitiveType_Cast) { - return false; - } - auto in_size = node->input_indices_.size(); - bool force_fp16 = false; - for (std::size_t k = 0; k < in_size; k++) { - schema::Tensor *tensor = model_->all_tensors_.at(node->input_indices_[k]); - if ((tensor->dataType() == kNumberTypeFloat16) && (tensor->nodeType() == NodeType_ValueNode)) { + schema::Tensor *tensor = model_.get()->all_tensors_.at(node->input_indices_[0]); + if (tensor->dataType() == kNumberTypeFloat16) { force_fp16 = true; - break; + } else if (tensor->dataType() == kNumberTypeFloat32) { + return false; + } + } else { + auto in_size = node->input_indices_.size(); + for (std::size_t k = 0; k < in_size; k++) { + schema::Tensor *tensor = model_->all_tensors_.at(node->input_indices_[k]); + if ((tensor->dataType() == kNumberTypeFloat16) && (tensor->nodeType() == NodeType_ValueNode)) { + force_fp16 = true; + break; + } } } + const auto &node_name = node->name_; bool is_fp16 = true; if (!force_fp16) { @@ -551,7 +548,7 @@ int TrainSession::RunGraph(const KernelCallBack &before, const KernelCallBack &a return lite::RET_NULL_PTR; } auto &run_kernels = (train_mode_) ? train_kernels_ : inference_kernels_; - if (context_->IsCpuFloat16Enabled() && !cfg_.mix_precision_cfg_.is_raw_mix_precision_) { + if (context_->IsCpuFloat16Enabled()) { ret = MixPrecisionExecKernels(before, after, run_kernels); } else { ret = ExecKernels(before, after, run_kernels); diff --git a/mindspore/lite/test/config/models_ms_train.cfg b/mindspore/lite/test/config/models_ms_train.cfg index fe887a5cde..add85f1d15 100644 --- a/mindspore/lite/test/config/models_ms_train.cfg +++ b/mindspore/lite/test/config/models_ms_train.cfg @@ -29,7 +29,7 @@ mini_alexnet fp16 6 nin fp16 8.0 #lenet fp16 2 mobilenetv1 fp16 2 -mobilenetv2 fp16 2 +mobilenetv2_mix_precision fp16 2 mobilenetv3 fp16 7 effnet fp16 2 effnet_tune fp16 5.0 diff --git a/mindspore/lite/test/ut/src/runtime/kernel/arm/cxx_api/model_test.cc b/mindspore/lite/test/ut/src/runtime/kernel/arm/cxx_api/model_test.cc index 2ea6bec5a1..167bcdb94a 100644 --- a/mindspore/lite/test/ut/src/runtime/kernel/arm/cxx_api/model_test.cc +++ b/mindspore/lite/test/ut/src/runtime/kernel/arm/cxx_api/model_test.cc @@ -192,19 +192,9 @@ TEST_F(TestCxxApiLiteModel, test_fp32_SUCCESS) { cpu_context->SetEnableFP16(true); context->MutableDeviceInfo().push_back(cpu_context); auto train_cfg = std::make_shared(); - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = true; ASSERT_TRUE(Serialization::Load("./nets/conv_train_model.ms", ModelType::kMindIR, &graph) == kSuccess); ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = false; - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - cpu_context->SetEnableFP16(false); - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = true; - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); } TEST_F(TestCxxApiLiteModel, test_fp16_SUCCESS) { @@ -215,19 +205,9 @@ TEST_F(TestCxxApiLiteModel, test_fp16_SUCCESS) { cpu_context->SetEnableFP16(true); context->MutableDeviceInfo().push_back(cpu_context); auto train_cfg = std::make_shared(); - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = true; ASSERT_TRUE(Serialization::Load("./nets/mix_lenet_tod.ms", ModelType::kMindIR, &graph) == kSuccess); ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = false; - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - cpu_context->SetEnableFP16(false); - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); - - train_cfg->mix_precision_cfg_.is_raw_mix_precision_ = true; - ASSERT_TRUE(model.Build(GraphCell(graph), context, train_cfg) == kSuccess); } #define NUM_OF_CLASSES 10 diff --git a/mindspore/lite/tools/benchmark_train/net_train.cc b/mindspore/lite/tools/benchmark_train/net_train.cc index dcced05209..0862b2d0cb 100644 --- a/mindspore/lite/tools/benchmark_train/net_train.cc +++ b/mindspore/lite/tools/benchmark_train/net_train.cc @@ -344,13 +344,10 @@ std::unique_ptr NetTrain::CreateAndRunNetworkForTrain(cons session::TrainSession::CreateTransferSession(bb_filename, filename, &context, true, &train_cfg)); if (session == nullptr) { MS_LOG(ERROR) << "RunNetTrain CreateTranferSession failed while running " << model_name.c_str(); - std::cout << "RunNetTrain CreateTranferSession failed while running " << model_name.c_str() << std::endl; return nullptr; } } else { MS_LOG(INFO) << "CreateTrainSession from model file" << filename.c_str(); - std::cout << "CreateTrainSession from model file " << filename.c_str() << std::endl; - std::cout << "Is raw mix precision model: " << train_cfg.mix_precision_cfg_.is_raw_mix_precision_ << std::endl; session = std::unique_ptr( session::TrainSession::CreateTrainSession(filename, &context, true, &train_cfg)); if (session == nullptr) { @@ -382,6 +379,7 @@ std::unique_ptr NetTrain::CreateAndRunNetworkForInference( auto *model = mindspore::lite::Model::Import(filenamems.c_str()); if (model == nullptr) { MS_LOG(ERROR) << "create model for train session failed"; + std::cout << "create model for train session failed " << filenamems.c_str() << std::endl; return nullptr; } session = std::unique_ptr(session::LiteSession::CreateSession(&context)); @@ -422,7 +420,6 @@ int NetTrain::CreateAndRunNetwork(const std::string &filename, const std::string } if (!(flags_->loss_name_.empty())) train_cfg.loss_name_.emplace_back(flags_->loss_name_); } - train_cfg.mix_precision_cfg_.is_raw_mix_precision_ = flags_->is_raw_mix_precision_; std::unique_ptr session; if (train_session) { session = CreateAndRunNetworkForTrain(filename, bb_filename, context, train_cfg, epochs); @@ -593,6 +590,7 @@ void NetTrain::CheckSum(mindspore::tensor::MSTensor *tensor, std::string node_ty #ifdef ENABLE_FP16 case kNumberTypeFloat16: std::cout << TensorSum(data, tensor_size) << std::endl; + TensorNan(reinterpret_cast(data), tensor_size); break; #endif default: diff --git a/mindspore/lite/tools/benchmark_train/net_train.h b/mindspore/lite/tools/benchmark_train/net_train.h index b6ad4fd26d..cffe37a041 100644 --- a/mindspore/lite/tools/benchmark_train/net_train.h +++ b/mindspore/lite/tools/benchmark_train/net_train.h @@ -31,11 +31,21 @@ #include #include +#ifdef ENABLE_FP16 +#include +#endif #include "tools/common/flag_parser.h" #include "src/common/file_utils.h" #include "src/common/utils.h" #include "include/lite_session.h" +#ifdef ENABLE_FP16 +static __attribute__((always_inline)) inline bool MS_ISNAN_FP16(float16_t var) { + volatile float16_t d = var; + return d != d; +} +#endif + namespace mindspore::lite { enum MS_API DataType { kImage = 0, kBinary = 1 }; @@ -75,8 +85,6 @@ class MS_API NetTrainFlags : public virtual FlagParser { AddFlag(&NetTrainFlags::virtual_batch_, "virtualBatch", "use virtual batch", false); AddFlag(&NetTrainFlags::resize_dims_in_, "inputShapes", "Shape of input data, the format should be NHWC. e.g. 1,32,32,32:1,1,32,32,1", ""); - AddFlag(&NetTrainFlags::is_raw_mix_precision_, "isRawMixPrecision", - "If model is mix precision export from MindSpore,please set true", false); } ~NetTrainFlags() override = default; @@ -109,7 +117,6 @@ class MS_API NetTrainFlags : public virtual FlagParser { std::vector> resize_dims_; std::string loss_name_ = ""; std::string inference_file_ = ""; - bool is_raw_mix_precision_ = false; }; class MS_API NetTrain { @@ -219,11 +226,21 @@ class MS_API NetTrain { void TensorNan(float *data, int size) { for (int i = 0; i < size; i++) { if (std::isnan(data[i])) { - std::cout << "nan value of index=" << i << std::endl; + std::cout << "nan value of index=" << i << ", " << data[i] << std::endl; break; } } } +#ifdef ENABLE_FP16 + void TensorNan(float16_t *data, int size) { + for (int i = 0; i < size; i++) { + if (MS_ISNAN_FP16(data[i]) || std::isinf(data[i])) { + std::cout << "nan or inf value of index=" << i << ", " << data[i] << std::endl; + break; + } + } + } +#endif NetTrainFlags *flags_; // callback parameters