forked from huawei/mindspore2022
fix dynamic infer shape bug
This commit is contained in:
parent
76891ef114
commit
8ba7731c9a
|
|
@ -20,6 +20,7 @@
|
|||
#include <unordered_map>
|
||||
#include <deque>
|
||||
#include <map>
|
||||
#include <tuple>
|
||||
#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<FuncGraphPtr> *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<CNodePtr> 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<FuncGraph>(input)) {
|
||||
continue;
|
||||
}
|
||||
auto new_fg = GetValueNode<FuncGraphPtr>(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<quant::FullQuantQuantizer>(*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<converter::Flags *>(config), &init_scale);
|
||||
if (status != RET_OK) {
|
||||
MS_LOG(ERROR) << "Grid search with scale failed.";
|
||||
return status;
|
||||
}
|
||||
auto quantizer = std::make_unique<quant::WeightQuantizer>(*config);
|
||||
if (quantizer == nullptr) {
|
||||
MS_LOG(ERROR) << "New WeightQuantizer failed";
|
||||
ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED);
|
||||
return RET_ERROR;
|
||||
}
|
||||
status = static_cast<quant::WeightQuantizer *>(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<quant::WeightQuantizer>(*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<quant::DynamicQuantizer>(*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<std::string, OpParameter *> 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<FuncGraphPtr> 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<converter::Flags *>(config));
|
||||
if (status != RET_OK) {
|
||||
MS_LOG(ERROR) << "Do Quantize failed.";
|
||||
return nullptr;
|
||||
|
|
|
|||
|
|
@ -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<FuncGraphPtr> *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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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<PrimitivePtr> support_weight_quant_nodes = {prim::kPrimMatMulFusion, prim::kPrimGather};
|
||||
const std::set<PrimitivePtr> 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<PrimitivePtr> 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
|
||||
|
|
|
|||
|
|
@ -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<PrimitivePtr> &support_dynamic_quant_ops) {
|
||||
MS_ASSERT(graph != nullptr);
|
||||
auto cnodes = graph->GetOrderedCnodes();
|
||||
const std::set<PrimitivePtr> support_dynamic_quant_ops = {
|
||||
prim::kPrimMatMulFusion,
|
||||
};
|
||||
for (auto &cnode : cnodes) {
|
||||
auto ret = CheckDataType(cnode, kNumberTypeFloat32);
|
||||
if (ret == RET_NO_CHANGE) {
|
||||
|
|
|
|||
|
|
@ -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 <vector>
|
||||
#include <set>
|
||||
#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<PrimitivePtr> &support_dynamic_quant_ops);
|
||||
|
||||
private:
|
||||
ValueNodePtr NewQuantCastValueNode(int src_type, int dst_type, const std::vector<schema::QuantParamT> &quant_params);
|
||||
|
|
|
|||
|
|
@ -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 <memory>
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <deque>
|
||||
#include <map>
|
||||
#include <set>
|
||||
#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<FuncGraphPtr> *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<CNodePtr> 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<FuncGraph>(input)) {
|
||||
continue;
|
||||
}
|
||||
auto new_fg = GetValueNode<FuncGraphPtr>(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<FullQuantQuantizer>(*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<converter::Flags *>(config), &init_scale);
|
||||
if (status != RET_OK) {
|
||||
MS_LOG(ERROR) << "Grid search with scale failed.";
|
||||
return status;
|
||||
}
|
||||
auto quantizer = std::make_unique<WeightQuantizer>(*config);
|
||||
if (quantizer == nullptr) {
|
||||
MS_LOG(ERROR) << "New WeightQuantizer failed";
|
||||
ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_MEMORY_FAILED);
|
||||
return RET_ERROR;
|
||||
}
|
||||
status = static_cast<WeightQuantizer *>(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<WeightQuantizer>(*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<DynamicQuantizer>(*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<std::string, OpParameter *> 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<FuncGraphPtr> 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
|
||||
|
|
@ -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 <utility>
|
||||
#include <map>
|
||||
#include <vector>
|
||||
#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
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue