From 8ba7731c9aa3073ee913f0b241b7f2805575dcfc Mon Sep 17 00:00:00 2001 From: yeyunpeng2020 Date: Sat, 15 Jan 2022 17:05:39 +0800 Subject: [PATCH] fix dynamic infer shape bug --- .../lite/tools/converter/anf_transform.cc | 195 +--------------- .../lite/tools/converter/anf_transform.h | 6 +- .../converter/quantizer/dynamic_quantizer.cc | 21 +- .../quantizer/insert_quant_node_manager.cc | 6 +- .../quantizer/insert_quant_node_manager.h | 3 +- .../quantizer/quantization_optimizer.cc | 214 ++++++++++++++++++ .../quantizer/quantization_optimizer.h | 36 +++ .../converter/quantizer/weight_quantizer.cc | 5 +- 8 files changed, 280 insertions(+), 206 deletions(-) create mode 100644 mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc create mode 100644 mindspore/lite/tools/converter/quantizer/quantization_optimizer.h diff --git a/mindspore/lite/tools/converter/anf_transform.cc b/mindspore/lite/tools/converter/anf_transform.cc index e73c9e0712..6f35b82db9 100644 --- a/mindspore/lite/tools/converter/anf_transform.cc +++ b/mindspore/lite/tools/converter/anf_transform.cc @@ -20,6 +20,7 @@ #include #include #include +#include #include "nnacl/op_base.h" #include "src/common/log_adapter.h" #include "tools/converter/optimizer_manager.h" @@ -74,9 +75,7 @@ #include "tools/optimizer/graph/special_node_postprocess.h" #include "tools/optimizer/graph/specify_graph_input_format.h" #include "tools/optimizer/graph/dump_graph.h" -#include "tools/converter/quantizer/full_quant_quantizer.h" -#include "tools/converter/quantizer/weight_quantizer.h" -#include "tools/converter/quantizer/dynamic_quantizer.h" +#include "tools/converter/quantizer/quantization_optimizer.h" #include "tools/optimizer/parallel/split_strategy.h" #include "tools/optimizer/parallel/spliter.h" #include "tools/optimizer/fisson/iter_node_outputs.h" @@ -88,8 +87,7 @@ #include "tools/optimizer/format/to_nchw_format.h" #include "tools/optimizer/format/to_nhwc_format.h" #include "tools/converter/adapter/acl/acl_pass.h" -#include "tools/converter/quantizer/parameter_tunner.h" -#include "tools/converter/quantizer/debug_info_manager.h" +#include "src/common/log_util.h" using std::string; namespace mindspore::lite { @@ -350,185 +348,12 @@ int AnfTransform::RunConstFoldPass(const FuncGraphPtr &old_graph, const converte return RET_OK; } -void AnfTransform::GetFuncGraphs(const FuncGraphPtr &func_graph, std::set *all_func_graphs) { - MS_ASSERT(func_graph != nullptr); - MS_ASSERT(all_func_graphs != nullptr); - all_func_graphs->insert(func_graph); - auto nodes = func_graph->GetOrderedCnodes(); - std::deque to_process{}; - to_process.insert(to_process.end(), nodes.begin(), nodes.end()); - while (!to_process.empty()) { - auto &cur_cnode = to_process.front(); - for (auto &input : cur_cnode->inputs()) { - if (!IsValueNode(input)) { - continue; - } - auto new_fg = GetValueNode(input); - if (all_func_graphs->find(new_fg) != all_func_graphs->end()) { - continue; - } - all_func_graphs->insert(new_fg); - auto new_nodes = new_fg->GetOrderedCnodes(); - to_process.insert(to_process.end(), new_nodes.begin(), new_nodes.end()); - } - to_process.pop_front(); - } -} - -int DoFullQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { - auto quantizer = std::make_unique(*config); - if (quantizer == nullptr) { - MS_LOG(ERROR) << "New FullQuantQuantizer failed"; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); - return RET_ERROR; - } - auto status = quantizer->DoQuantize(old_graph); - if (status != RET_OK) { - MS_LOG(ERROR) << "DoQuantization failed " << status; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); - return RET_ERROR; - } - return RET_OK; -} - -int DoWeightQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { - double init_scale = config->mixedBitWeightQuantParam.init_scale; - if (config->commonQuantParam.bit_num == 0 && config->mixedBitWeightQuantParam.auto_tune) { - quant::ParameterOptimizer optimizer; - auto status = optimizer.GridSearchForScale(old_graph, const_cast(config), &init_scale); - if (status != RET_OK) { - MS_LOG(ERROR) << "Grid search with scale failed."; - return status; - } - auto quantizer = std::make_unique(*config); - if (quantizer == nullptr) { - MS_LOG(ERROR) << "New WeightQuantizer failed"; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); - return RET_ERROR; - } - status = static_cast(quantizer.get())->DoQuantize(old_graph, init_scale); - if (status != RET_OK) { - MS_LOG(ERROR) << "DoQuantization failed " << status; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); - return RET_ERROR; - } - } else { - auto quantizer = std::make_unique(*config); - if (quantizer == nullptr) { - MS_LOG(ERROR) << "New WeightQuantizer failed"; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); - return RET_ERROR; - } - auto status = quantizer->DoQuantize(old_graph); - if (status != RET_OK) { - MS_LOG(ERROR) << "DoQuantization failed " << status; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); - return RET_ERROR; - } - } - return RET_OK; -} - -int DoDynamicQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { - auto quantizer = std::make_unique(*config); - if (quantizer == nullptr) { - MS_LOG(ERROR) << "New DynamicQuantizer failed"; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); - return RET_ERROR; - } - auto status = quantizer->DoQuantize(old_graph); - if (status != RET_OK) { - MS_LOG(ERROR) << "DoQuantization failed " << status; - ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); - return RET_ERROR; - } - return RET_OK; -} - -int DoQuantDebug(const FuncGraphPtr &old_graph, const converter::Flags *config, const quant::SessionModel &origin) { - auto quant = quant::CreateSessionByFuncGraph(old_graph, *config, config->commonQuantParam.thread_num); - std::map op_parameters; - FetchOpParameterFromFuncGraph(old_graph, &op_parameters); - DebugInfoManager manager; - CHECK_NULL_RETURN(origin.model); - CHECK_NULL_RETURN(origin.session); - CHECK_NULL_RETURN(quant.model); - CHECK_NULL_RETURN(quant.session); - auto status = manager.CompareOriginWithQuant( - origin, quant, op_parameters, config->commonQuantParam.debug_info_save_path, config->dataPreProcessParam); - auto free_buffer = [&] { - delete origin.session; - delete origin.model; - delete quant.session; - delete quant.model; - for (auto parameter : op_parameters) { - if (parameter.second != nullptr) { - free(parameter.second); - parameter.second = nullptr; - } - } - op_parameters.clear(); - }; - if (status != RET_OK) { - MS_LOG(ERROR) << "Compare origin with quant failed."; - free_buffer(); - return status; - } - free_buffer(); - return RET_OK; -} - -int AnfTransform::DoSingleGraphQuantize(const FuncGraphPtr &old_graph, const converter::Flags *config) { - // quant - if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_NONE) { - return RET_OK; - } - int status; - - quant::SessionModel origin; - if (config->commonQuantParam.is_debug) { - converter::Flags new_flag = *config; - new_flag.commonQuantParam.quant_type = schema::QuantType_QUANT_NONE; - origin = quant::CreateSessionByFuncGraph(old_graph, new_flag, config->commonQuantParam.thread_num); - } - if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_ALL) { - status = DoFullQuant(old_graph, config); - if (status != RET_OK) { - MS_LOG(ERROR) << "Do full quant failed."; - return status; - } - } else if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_WEIGHT) { - status = DoWeightQuant(old_graph, config); - if (status != RET_OK) { - MS_LOG(ERROR) << "Do weight quant failed."; - return status; - } - } else if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_DANAMIC) { - status = DoDynamicQuant(old_graph, config); - if (status != RET_OK) { - MS_LOG(ERROR) << "Do dynamic quant failed."; - return status; - } - } - if (config->commonQuantParam.is_debug) { - status = DoQuantDebug(old_graph, config, origin); - if (status != RET_OK) { - MS_LOG(ERROR) << "Do quant debug failed."; - return status; - } - } - return RET_OK; -} - -int AnfTransform::DoQuantize(const FuncGraphPtr &old_graph, const converter::Flags *config) { - std::set all_func_graphs{}; - GetFuncGraphs(old_graph, &all_func_graphs); - for (auto &item : all_func_graphs) { - auto status = DoSingleGraphQuantize(item, config); - if (status != RET_OK) { - MS_LOG(ERROR) << "Do Quantize failed."; - return status; - } +int AnfTransform::DoQuantize(const FuncGraphPtr &old_graph, converter::Flags *config) { + quant::QuantizationOptimizer optimizer(config); + auto ret = optimizer.Run(old_graph); + if (ret != RET_OK) { + MS_LOG(ERROR) << "Post training quantization failed."; + return ret; } return RET_OK; } @@ -619,7 +444,7 @@ FuncGraphPtr AnfTransform::TransformFuncGraph(const FuncGraphPtr &old_graph, con MS_LOG(ERROR) << "Unsupported external extension with quantization."; return nullptr; } - status = DoQuantize(old_graph, config); + status = DoQuantize(old_graph, const_cast(config)); if (status != RET_OK) { MS_LOG(ERROR) << "Do Quantize failed."; return nullptr; diff --git a/mindspore/lite/tools/converter/anf_transform.h b/mindspore/lite/tools/converter/anf_transform.h index 2648f6bb53..fa9b567447 100644 --- a/mindspore/lite/tools/converter/anf_transform.h +++ b/mindspore/lite/tools/converter/anf_transform.h @@ -49,11 +49,7 @@ class AnfTransform { static int RunParallelPass(const FuncGraphPtr &old_graph, const converter::Flags *config); - int DoQuantize(const FuncGraphPtr &old_graph, const converter::Flags *config); - - static void GetFuncGraphs(const FuncGraphPtr &func_graph, std::set *all_func_graphs); - - int DoSingleGraphQuantize(const FuncGraphPtr &old_graph, const converter::Flags *config); + static int DoQuantize(const FuncGraphPtr &old_graph, converter::Flags *config); static bool StoreBuiltinPass(const converter::Flags *config); diff --git a/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc b/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc index 31f3895e8b..2d027078b2 100644 --- a/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc +++ b/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc @@ -20,22 +20,27 @@ namespace mindspore::lite::quant { int DynamicQuantizer::DoQuantize(FuncGraphPtr func_graph) { - InsertQuantNodeManager manager; - auto ret = manager.InsertDynamicQuantPass(func_graph); - if (ret != RET_OK) { - MS_LOG(ERROR) << "Insert dynamic quant failed."; - return ret; - } - auto quantizer = WeightQuantizer(flags_); + // Dynamic dont support filters. flags_.commonQuantParam.min_quant_weight_channel = 0; flags_.commonQuantParam.min_quant_weight_size = 0; + flags_.commonQuantParam.skip_quant_node.clear(); + auto quantizer = WeightQuantizer(flags_); const std::set support_weight_quant_nodes = {prim::kPrimMatMulFusion, prim::kPrimGather}; const std::set symmetric_nodes = {prim::kPrimMatMulFusion}; - ret = quantizer.WeightQuant(func_graph, support_weight_quant_nodes, {}, symmetric_nodes); + auto ret = quantizer.WeightQuant(func_graph, support_weight_quant_nodes, {}, symmetric_nodes); if (ret != RET_OK) { MS_LOG(ERROR) << "Weight Quant failed."; return ret; } + InsertQuantNodeManager manager; + const std::set support_dynamic_quant_ops = { + prim::kPrimMatMulFusion, + }; + ret = manager.InsertDynamicQuantPass(func_graph, support_dynamic_quant_ops); + if (ret != RET_OK) { + MS_LOG(ERROR) << "Insert dynamic quant failed."; + return ret; + } return RET_OK; } } // namespace mindspore::lite::quant diff --git a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc index ad8cab44c5..04733cdf1c 100644 --- a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc +++ b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc @@ -198,12 +198,10 @@ int InsertQuantNodeManager::MarkDynamicQuantize(const CNodePtr &cnode) { return RET_OK; } -int InsertQuantNodeManager::InsertDynamicQuantPass(const FuncGraphPtr &graph) { +int InsertQuantNodeManager::InsertDynamicQuantPass(const FuncGraphPtr &graph, + const std::set &support_dynamic_quant_ops) { MS_ASSERT(graph != nullptr); auto cnodes = graph->GetOrderedCnodes(); - const std::set support_dynamic_quant_ops = { - prim::kPrimMatMulFusion, - }; for (auto &cnode : cnodes) { auto ret = CheckDataType(cnode, kNumberTypeFloat32); if (ret == RET_NO_CHANGE) { diff --git a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.h b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.h index e1390b3166..6767d875d9 100644 --- a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.h +++ b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.h @@ -17,6 +17,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_INSERT_QUANT_NODE_MANAGER_H #define MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_INSERT_QUANT_NODE_MANAGER_H #include +#include #include "include/errorcode.h" #include "ir/anf.h" #include "ir/dtype/type_id.h" @@ -32,7 +33,7 @@ class InsertQuantNodeManager { int InsertQuantDtypeCastPass(const FuncGraphPtr &graph); - int InsertDynamicQuantPass(const FuncGraphPtr &graph); + int InsertDynamicQuantPass(const FuncGraphPtr &graph, const std::set &support_dynamic_quant_ops); private: ValueNodePtr NewQuantCastValueNode(int src_type, int dst_type, const std::vector &quant_params); diff --git a/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc new file mode 100644 index 0000000000..5dd0275579 --- /dev/null +++ b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc @@ -0,0 +1,214 @@ +/** + * Copyright 2022 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "tools/converter/quantizer/quantization_optimizer.h" +#include +#include +#include +#include +#include +#include +#include "tools/anf_exporter/fetch_content.h" +#include "base/base.h" +#include "tools/converter/quantizer/quantize_util.h" +#include "tools/converter/quantizer/weight_quantizer.h" +#include "tools/converter/quantizer/full_quant_quantizer.h" +#include "tools/converter/quantizer/debug_info_manager.h" +#include "tools/converter/quantizer/parameter_tunner.h" +#include "tools/converter/quantizer/dynamic_quantizer.h" + +namespace mindspore::lite::quant { +void GetFuncGraphs(const FuncGraphPtr &func_graph, std::set *all_func_graphs) { + MS_ASSERT(func_graph != nullptr); + MS_ASSERT(all_func_graphs != nullptr); + all_func_graphs->insert(func_graph); + auto nodes = func_graph->GetOrderedCnodes(); + std::deque to_process{}; + to_process.insert(to_process.end(), nodes.begin(), nodes.end()); + while (!to_process.empty()) { + auto &cur_cnode = to_process.front(); + for (auto &input : cur_cnode->inputs()) { + if (!IsValueNode(input)) { + continue; + } + auto new_fg = GetValueNode(input); + if (all_func_graphs->find(new_fg) != all_func_graphs->end()) { + continue; + } + all_func_graphs->insert(new_fg); + auto new_nodes = new_fg->GetOrderedCnodes(); + to_process.insert(to_process.end(), new_nodes.begin(), new_nodes.end()); + } + to_process.pop_front(); + } +} + +int DoFullQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { + auto quantizer = std::make_unique(*config); + if (quantizer == nullptr) { + MS_LOG(ERROR) << "New FullQuantQuantizer failed"; + return RET_ERROR; + } + auto status = quantizer->DoQuantize(old_graph); + if (status != RET_OK) { + MS_LOG(ERROR) << "DoQuantization failed " << status; + return RET_ERROR; + } + return RET_OK; +} + +int DoWeightQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { + double init_scale = config->mixedBitWeightQuantParam.init_scale; + if (config->commonQuantParam.bit_num == 0 && config->mixedBitWeightQuantParam.auto_tune) { + ParameterOptimizer optimizer; + auto status = optimizer.GridSearchForScale(old_graph, const_cast(config), &init_scale); + if (status != RET_OK) { + MS_LOG(ERROR) << "Grid search with scale failed."; + return status; + } + auto quantizer = std::make_unique(*config); + if (quantizer == nullptr) { + MS_LOG(ERROR) << "New WeightQuantizer failed"; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); + return RET_ERROR; + } + status = static_cast(quantizer.get())->DoQuantize(old_graph, init_scale); + if (status != RET_OK) { + MS_LOG(ERROR) << "DoQuantization failed " << status; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); + return RET_ERROR; + } + } else { + auto quantizer = std::make_unique(*config); + if (quantizer == nullptr) { + MS_LOG(ERROR) << "New WeightQuantizer failed"; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); + return RET_ERROR; + } + auto status = quantizer->DoQuantize(old_graph); + if (status != RET_OK) { + MS_LOG(ERROR) << "DoQuantization failed " << status; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); + return RET_ERROR; + } + } + return RET_OK; +} + +int DoDynamicQuant(const FuncGraphPtr &old_graph, const converter::Flags *config) { + auto quantizer = std::make_unique(*config); + if (quantizer == nullptr) { + MS_LOG(ERROR) << "New DynamicQuantizer failed"; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED); + return RET_ERROR; + } + auto status = quantizer->DoQuantize(old_graph); + if (status != RET_OK) { + MS_LOG(ERROR) << "DoQuantization failed " << status; + ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); + return RET_ERROR; + } + return RET_OK; +} + +int DoQuantDebug(const FuncGraphPtr &old_graph, const converter::Flags *config, const SessionModel &origin) { + auto quant = CreateSessionByFuncGraph(old_graph, *config, config->commonQuantParam.thread_num); + std::map op_parameters; + FetchOpParameterFromFuncGraph(old_graph, &op_parameters); + DebugInfoManager manager; + CHECK_NULL_RETURN(origin.model); + CHECK_NULL_RETURN(origin.session); + CHECK_NULL_RETURN(quant.model); + CHECK_NULL_RETURN(quant.session); + auto status = manager.CompareOriginWithQuant( + origin, quant, op_parameters, config->commonQuantParam.debug_info_save_path, config->dataPreProcessParam); + auto free_buffer = [&] { + delete origin.session; + delete origin.model; + delete quant.session; + delete quant.model; + for (auto parameter : op_parameters) { + if (parameter.second != nullptr) { + free(parameter.second); + parameter.second = nullptr; + } + } + op_parameters.clear(); + }; + if (status != RET_OK) { + MS_LOG(ERROR) << "Compare origin with quant failed."; + free_buffer(); + return status; + } + free_buffer(); + return RET_OK; +} + +int DoSingleGraphQuantize(const FuncGraphPtr &old_graph, const converter::Flags *config) { + // quant + if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_NONE) { + return RET_OK; + } + int status; + + SessionModel origin; + if (config->commonQuantParam.is_debug) { + converter::Flags new_flag = *config; + new_flag.commonQuantParam.quant_type = schema::QuantType_QUANT_NONE; + origin = CreateSessionByFuncGraph(old_graph, new_flag, config->commonQuantParam.thread_num); + } + if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_ALL) { + status = DoFullQuant(old_graph, config); + if (status != RET_OK) { + MS_LOG(ERROR) << "Do full quant failed."; + return status; + } + } else if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_WEIGHT) { + status = DoWeightQuant(old_graph, config); + if (status != RET_OK) { + MS_LOG(ERROR) << "Do weight quant failed."; + return status; + } + } else if (config->commonQuantParam.quant_type == schema::QuantType_QUANT_DANAMIC) { + status = DoDynamicQuant(old_graph, config); + if (status != RET_OK) { + MS_LOG(ERROR) << "Do dynamic quant failed."; + return status; + } + } + if (config->commonQuantParam.is_debug) { + status = DoQuantDebug(old_graph, config, origin); + if (status != RET_OK) { + MS_LOG(ERROR) << "Do quant debug failed."; + return status; + } + } + return RET_OK; +} + +int QuantizationOptimizer::Run(const mindspore::FuncGraphPtr &func_graph) { + std::set all_func_graphs{}; + GetFuncGraphs(func_graph, &all_func_graphs); + for (auto &item : all_func_graphs) { + auto status = DoSingleGraphQuantize(item, flags_); + if (status != RET_OK) { + MS_LOG(ERROR) << "Do Quantize failed."; + return status; + } + } + return RET_OK; +} +} // namespace mindspore::lite::quant diff --git a/mindspore/lite/tools/converter/quantizer/quantization_optimizer.h b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.h new file mode 100644 index 0000000000..f71947bfed --- /dev/null +++ b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.h @@ -0,0 +1,36 @@ +/** + * Copyright 2022 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_QUANTIZATION_OPTIMIZER_H +#define MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_QUANTIZATION_OPTIMIZER_H +#include +#include +#include +#include "backend/optimizer/common/pass.h" +#include "tools/converter/converter_flags.h" + +namespace mindspore::lite::quant { +class QuantizationOptimizer { + public: + explicit QuantizationOptimizer(converter::Flags *flags) : flags_(flags) {} + ~QuantizationOptimizer() = default; + int Run(const FuncGraphPtr &func_graph); + + private: + converter::Flags *flags_; +}; +} // namespace mindspore::lite::quant +#endif // MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_QUANTIZATION_OPTIMIZER_H diff --git a/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc b/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc index ba543c8cbc..614ab0d88d 100644 --- a/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc +++ b/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc @@ -143,9 +143,8 @@ int WeightQuantizer::DoMarkWeightQuantizeIfQuantized(const CNodePtr &cnode) { } auto quant_param_holder = GetCNodeQuantHolder(primitive); - if (quant_param_holder->quant_type() == schema::QuantType_QUANT_WEIGHT || - quant_param_holder->quant_type() == schema::QuantType_QUANT_DANAMIC) { - // already marked with QuantType_QUANT_WEIGHT or QuantType_QUANT_DANAMIC + if (quant_param_holder->quant_type() == schema::QuantType_QUANT_WEIGHT) { + // already marked with QuantType_QUANT_WEIGHT return RET_OK; }