forked from huawei/mindspore2022
reduce inited in quant_params
This commit is contained in:
parent
7fbc6d6865
commit
6fd4f64569
|
|
@ -100,16 +100,21 @@ int CoderGraph::ConvertTensors() {
|
|||
auto quant_params = origin_tensor->quantParams();
|
||||
if (quant_params != nullptr) {
|
||||
for (int j = 0; j < static_cast<int>(quant_params->size()); j++) {
|
||||
auto quant_param = quant_params->Get(j);
|
||||
LiteQuantParam quant_arg{};
|
||||
quant_arg.bitNum = quant_params->Get(j)->numBits();
|
||||
quant_arg.scale = quant_params->Get(j)->scale();
|
||||
quant_arg.zeroPoint = quant_params->Get(j)->zeroPoint();
|
||||
quant_arg.var_corr = quant_params->Get(j)->varCorr();
|
||||
quant_arg.mean_corr = quant_params->Get(j)->meanCorr();
|
||||
quant_arg.inited = quant_params->Get(j)->inited();
|
||||
quant_arg.roundType = quant_params->Get(j)->roundType();
|
||||
quant_arg.multiplier = quant_params->Get(j)->multiplier();
|
||||
quant_arg.dstDtype = quant_params->Get(j)->dstDtype();
|
||||
if (quant_param == nullptr) {
|
||||
quant_arg.inited = false;
|
||||
} else {
|
||||
quant_arg.inited = true;
|
||||
quant_arg.bitNum = quant_param->numBits();
|
||||
quant_arg.scale = quant_param->scale();
|
||||
quant_arg.zeroPoint = quant_param->zeroPoint();
|
||||
quant_arg.var_corr = quant_param->varCorr();
|
||||
quant_arg.mean_corr = quant_param->meanCorr();
|
||||
quant_arg.roundType = quant_param->roundType();
|
||||
quant_arg.multiplier = quant_param->multiplier();
|
||||
quant_arg.dstDtype = quant_param->dstDtype();
|
||||
}
|
||||
dstTensor->AddQuantParam(quant_arg);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -54,7 +54,7 @@ bool NeedBitUppackCheck(const schema::Tensor &src_tensor) {
|
|||
return true;
|
||||
}
|
||||
bool need_bit_unpack = src_tensor.quantParams() != nullptr && src_tensor.quantParams()->size() > 0 &&
|
||||
src_tensor.quantParams()->Get(0) != nullptr && src_tensor.quantParams()->Get(0)->inited();
|
||||
src_tensor.quantParams()->Get(0) != nullptr;
|
||||
if (need_bit_unpack) {
|
||||
auto num_bits = src_tensor.quantParams()->Get(0)->numBits();
|
||||
need_bit_unpack = ((num_bits >= kBitNum1 && num_bits < kBitNum8) || (num_bits > kBitNum8 && num_bits < kBitNum16));
|
||||
|
|
@ -100,16 +100,21 @@ void LiteSession::ConvertTensorsQuantParam(const schema::Tensor *src_tensor, lit
|
|||
auto quant_params = src_tensor->quantParams();
|
||||
if (quant_params != nullptr) {
|
||||
for (size_t j = 0; j < quant_params->size(); j++) {
|
||||
auto quant_param = quant_params->Get(j);
|
||||
LiteQuantParam quant_arg{};
|
||||
quant_arg.bitNum = quant_params->Get(j)->numBits();
|
||||
quant_arg.scale = quant_params->Get(j)->scale();
|
||||
quant_arg.zeroPoint = quant_params->Get(j)->zeroPoint();
|
||||
quant_arg.var_corr = quant_params->Get(j)->varCorr();
|
||||
quant_arg.mean_corr = quant_params->Get(j)->meanCorr();
|
||||
quant_arg.inited = quant_params->Get(j)->inited();
|
||||
quant_arg.roundType = quant_params->Get(j)->roundType();
|
||||
quant_arg.multiplier = quant_params->Get(j)->multiplier();
|
||||
quant_arg.dstDtype = quant_params->Get(j)->dstDtype();
|
||||
if (quant_param == nullptr) {
|
||||
quant_arg.inited = false;
|
||||
} else {
|
||||
quant_arg.inited = true;
|
||||
quant_arg.bitNum = quant_param->numBits();
|
||||
quant_arg.scale = quant_param->scale();
|
||||
quant_arg.zeroPoint = quant_param->zeroPoint();
|
||||
quant_arg.var_corr = quant_param->varCorr();
|
||||
quant_arg.mean_corr = quant_param->meanCorr();
|
||||
quant_arg.roundType = quant_param->roundType();
|
||||
quant_arg.multiplier = quant_param->multiplier();
|
||||
quant_arg.dstDtype = quant_param->dstDtype();
|
||||
}
|
||||
dst_tensor->AddQuantParam(quant_arg);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -277,7 +277,7 @@ int WeightDecoder::UnPackToInt(const schema::Tensor &src_tensor, lite::Tensor *d
|
|||
return RET_NO_CHANGE;
|
||||
}
|
||||
auto quant_param = quant_params->Get(0);
|
||||
if (quant_param == nullptr || !quant_param->inited()) {
|
||||
if (quant_param == nullptr) {
|
||||
return RET_NO_CHANGE;
|
||||
}
|
||||
auto dst_data = dst_tensor->data();
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@
|
|||
* limitations under the License.
|
||||
*/
|
||||
#include "tools/converter/legacy_optimizer/graph/set_unused_quant_param_to_default_pass.h"
|
||||
#include "tools/converter/converter_context.h"
|
||||
#include "tools/common/tensor_util.h"
|
||||
#include "src/common/log_util.h"
|
||||
|
||||
|
|
@ -22,10 +21,20 @@ namespace mindspore::lite {
|
|||
STATUS SetUnusedQuantParamToDefaultPass::Run(schema::MetaGraphT *graph) {
|
||||
CHECK_NULL_RETURN(graph);
|
||||
for (auto &tensor : graph->allTensors) {
|
||||
bool has_quant_param = false;
|
||||
for (auto &quant_param : tensor->quantParams) {
|
||||
quant_param->min = 0;
|
||||
quant_param->max = 0;
|
||||
quant_param->min = 0.0;
|
||||
quant_param->max = 0.0;
|
||||
quant_param->narrowRange = true;
|
||||
if (quant_param->inited) {
|
||||
has_quant_param = true;
|
||||
quant_param->inited = false;
|
||||
} else {
|
||||
quant_param = nullptr;
|
||||
}
|
||||
}
|
||||
if (!has_quant_param) {
|
||||
tensor->quantParams.clear();
|
||||
}
|
||||
}
|
||||
return RET_OK;
|
||||
|
|
|
|||
Loading…
Reference in New Issue