From 87174d9561f1b39fa7eae2155add0853cb893518 Mon Sep 17 00:00:00 2001 From: wang_shaocong Date: Fri, 10 Sep 2021 17:27:20 +0800 Subject: [PATCH] [MSLITE] Convert a tf model whose input is also an output. --- mindspore/lite/src/lite_session.cc | 14 ++++----- mindspore/lite/src/mindrt_executor.cc | 29 +++++++++++++------ .../converter/parser/tf/tf_model_parser.cc | 5 +++- 3 files changed, 31 insertions(+), 17 deletions(-) diff --git a/mindspore/lite/src/lite_session.cc b/mindspore/lite/src/lite_session.cc index 0caaaffe44..b26620e197 100644 --- a/mindspore/lite/src/lite_session.cc +++ b/mindspore/lite/src/lite_session.cc @@ -209,7 +209,10 @@ int LiteSession::ConvertTensors(const lite::Model *model) { MS_LOG(ERROR) << "Graph output shouldn't have data"; return RET_ERROR; } - dst_tensor->set_category(Tensor::GRAPH_OUTPUT); + // a tensor is as both input and output, would be treated as an input. + if (!dst_tensor->IsGraphInput()) { + dst_tensor->set_category(Tensor::GRAPH_OUTPUT); + } } if (src_tensor->name() != nullptr) { dst_tensor->set_tensor_name(src_tensor->name()->str()); @@ -364,16 +367,13 @@ void LiteSession::InitGraphInOutTensorsMap(const lite::Model *model) { InitGraphInputMap(model); InitGraphOutputNodeMap(model); InitGraphOutputTensorMap(model); - for (auto *tensor : this->inputs_) { - tensor->set_category(Tensor::Category::GRAPH_INPUT); - } - for (auto *tensor : this->outputs_) { - tensor->set_category(Tensor::Category::GRAPH_OUTPUT); - } } void LiteSession::IsolateOutputTensor() { for (Tensor *src_tensor : outputs_) { + if (src_tensor->IsGraphInput()) { + continue; + } Tensor *new_tensor = new Tensor(src_tensor->data_type(), src_tensor->shape(), src_tensor->format(), Tensor::GRAPH_OUTPUT); new_tensor->set_allocator(src_tensor->allocator()); /* GPU use opencl allocator */ diff --git a/mindspore/lite/src/mindrt_executor.cc b/mindspore/lite/src/mindrt_executor.cc index f6841f78d1..65abfd5fc7 100644 --- a/mindspore/lite/src/mindrt_executor.cc +++ b/mindspore/lite/src/mindrt_executor.cc @@ -24,16 +24,24 @@ namespace mindspore::lite { void MindrtExecutor::PrepareInputData(const std::vector &kernels, const std::vector &inputs) { - for (size_t i = 0; i < inputs.size(); ++i) { - for (size_t j = 0; j < kernels.size(); ++j) { - auto in_tensor_size = kernels[j]->in_tensors().size(); - for (size_t k = 0; k < in_tensor_size; ++k) { - if (inputs[i] != kernels[j]->in_tensors()[k]) { - continue; - } - auto data = std::make_shared>(op_actors_[j]->GetAID(), inputs[i], static_cast(k)); - input_data_.emplace_back(data); + for (size_t j = 0; j < kernels.size(); ++j) { + auto in_tensor_size = kernels[j]->in_tensors().size(); + for (size_t k = 0; k < in_tensor_size; ++k) { + auto tensor = kernels[j]->in_tensors()[k]; + if (!tensor->IsGraphInput()) { + continue; } + size_t idx = std::find(inputs.begin(), inputs.end(), tensor) - inputs.begin(); + if (idx == inputs.size()) { + MS_LOG(ERROR) << "The input is not found."; + return; + } + auto data = std::make_shared>(op_actors_[j]->GetAID(), inputs.at(idx), static_cast(k)); + if (data == nullptr) { + MS_LOG(ERROR) << "new opdata failed."; + return; + } + input_data_.emplace_back(data); } } } @@ -42,6 +50,9 @@ void MindrtExecutor::PrepareOutputData(const std::vector & const std::vector &outputs) { for (size_t i = 0; i < outputs.size(); ++i) { Tensor *graph_output_tensor = outputs[i]; + if (graph_output_tensor->IsGraphInput()) { + continue; + } auto current_output_map = std::find_if(output_tensor_map_->begin(), output_tensor_map_->end(), [&](const auto output_map_tensor) { if (graph_output_tensor == output_map_tensor.second) { diff --git a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc index 0225b511d1..ee4160d09c 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc @@ -1059,7 +1059,10 @@ STATUS TFModelParser::ConvertOps(const tensorflow::NodeDef &node_def, MS_ASSERT(func_graph_ptr != nullptr); STATUS status = RET_OK; const auto &op_type = node_def.op(); - if (op_type == "Placeholder" || op_type == "Const" || op_type == "Identity" || op_type == "StopGradient") { + if (op_type == "Identity" || op_type == "StopGradient") { + return RET_OK; + } else if (op_type == "Placeholder" || op_type == "Const") { + node_output_num_[node_def.name()] = 1; return RET_OK; }