From 401084e0a68fe210e06dac521799eeac9717310e Mon Sep 17 00:00:00 2001 From: Alexander Malyshev Date: Thu, 24 Feb 2022 12:30:56 +0700 Subject: [PATCH] ONNX GNMT: add support features to converter --- .../transform/express_ir/onnx_exporter.cc | 903 ++++++++++-------- 1 file changed, 494 insertions(+), 409 deletions(-) diff --git a/mindspore/ccsrc/transform/express_ir/onnx_exporter.cc b/mindspore/ccsrc/transform/express_ir/onnx_exporter.cc index bba01693403..1d61a7e73dd 100644 --- a/mindspore/ccsrc/transform/express_ir/onnx_exporter.cc +++ b/mindspore/ccsrc/transform/express_ir/onnx_exporter.cc @@ -14,20 +14,20 @@ * limitations under the License. */ +#include +#include #include #include -#include #include -#include -#include +#include -#include "utils/hash_map.h" -#include "ir/tensor.h" -#include "ir/param_info.h" -#include "ir/func_graph.h" #include "base/core_ops.h" +#include "ir/func_graph.h" +#include "ir/param_info.h" +#include "ir/tensor.h" #include "proto/onnx.pb.h" #include "utils/check_convert_utils.h" +#include "utils/hash_map.h" #include "utils/ms_context.h" namespace mindspore { @@ -37,6 +37,7 @@ const int kOneNum = 1; const int kTwoNum = 2; const int kThreeNum = 3; const int kFourNum = 4; +const int kFiveNum = 5; const int64_t kOneNumLong = 1; const float weight_for_mul = 0.5; enum OpMergeMode { @@ -62,6 +63,12 @@ bool IsIgnoredIdentityNode(const AnfNodePtr &node) { return IsPrimitiveCNode(node, prim::kPrimDepend) || IsPrimitiveCNode(node, prim::kPrimLoad); } +/* + If true, the node should not be referenced by anything and should not be contributing to any + ref counts itself + */ +bool IsZeroRefcountNode(const AnfNodePtr &node) { return HasAbstractMonad(node) || IsIgnoredIdentityNode(node); } + // Ideally this should be applied to every node->input() call, not only inside GetNodeInputName static AnfNodePtr GetRealInput(const AnfNodePtr &origin_input) { AnfNodePtr input = origin_input; @@ -198,6 +205,21 @@ std::string MakeOutputName(const std::string &node_name, int output_index) { return node_name + "_" + std::to_string(output_index); } +size_t RavelIndex(const std::vector &index, const std::vector &shape) { + MS_EXCEPTION_IF_CHECK_FAIL(index.size() <= shape.size(), "Index ndims must be <= shape ndims"); + size_t result = 0; + size_t stride = 1; + for (size_t i = 0; i < shape.size() - index.size(); ++i) { + stride *= shape[shape.size() - 1 - i]; + } + for (size_t i = 0; i < index.size(); ++i) { + size_t rev_i = index.size() - 1 - i; + result += index[rev_i] * stride; + stride *= shape[rev_i]; + } + return result; +} + namespace fp16 { uint32_t FieldMask(unsigned int field_size) { const unsigned int BYTE_SIZE = 8; @@ -723,7 +745,6 @@ OPERATOR_ONNX_CONVERT_DEFINE(Minimum, Min, .CastInput(1, onnx::TensorProto_DataType_INT32, onnx::TensorProto_DataType_FLOAT) .CastOutputToInputType(0)) OPERATOR_ONNX_CONVERT_DEFINE(Transpose, Transpose, OpNameInfo()) -OPERATOR_ONNX_CONVERT_DEFINE(StridedSlice, Slice, OpNameInfo()) OPERATOR_ONNX_CONVERT_DEFINE(Exp, Exp, OpNameInfo()) OPERATOR_ONNX_CONVERT_DEFINE(Softplus, Softplus, OpNameInfo()) OPERATOR_ONNX_CONVERT_DEFINE(Tanh, Tanh, OpNameInfo()) @@ -759,6 +780,10 @@ OPERATOR_ONNX_CONVERT_DEFINE(ReverseSequence, ReverseSequence, .Attr("batch_dim", "batch_axis", onnx::AttributeProto_AttributeType_INT, SetAttrValueToProto) .CastInput(1, onnx::TensorProto_DataType_INT32, onnx::TensorProto_DataType_INT64)) +OPERATOR_ONNX_CONVERT_DEFINE(Less, Less, OpNameInfo()) +OPERATOR_ONNX_CONVERT_DEFINE(TensorScatterUpdate, ScatterND, + OpNameInfo().CastInput(1, onnx::TensorProto_DataType_INT32, + onnx::TensorProto_DataType_INT64)) #define OP_CONVERT_FUNCTION_NAME(name) GetOpOnnxConvertInfo_##name @@ -801,9 +826,11 @@ void RegisterOpConverters(const std::function &fn) { fn(OP_CONVERT_FUNCTION_NAME(GatherNd)()); fn(OP_CONVERT_FUNCTION_NAME(Select)()); fn(OP_CONVERT_FUNCTION_NAME(Log)()); + fn(OP_CONVERT_FUNCTION_NAME(Less)()); fn(OP_CONVERT_FUNCTION_NAME(Greater)()); fn(OP_CONVERT_FUNCTION_NAME(LogicalAnd)()); fn(OP_CONVERT_FUNCTION_NAME(ReverseSequence)()); + fn(OP_CONVERT_FUNCTION_NAME(TensorScatterUpdate)()); } class OpConvertRegistry { @@ -839,12 +866,14 @@ class OnnxExporter { private: void InitModelInfo(); - void ExportFuncGraph(const FuncGraphPtr &func_graph, onnx::GraphProto *graph_proto); - void ExportParameters(const FuncGraphPtr &func_graph, onnx::GraphProto *graph_proto); + void ExportFuncGraph(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + onnx::GraphProto *graph_proto, bool export_inputs = true); + void ExportInputs(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + onnx::GraphProto *graph_proto); - size_t ExportPrimitive(const FuncGraphPtr &func_graph, std::map *node_map_ptr, - const PrimitivePtr &prim, const std::vector &inputs, - onnx::GraphProto *graph_proto); + std::string ExportPrimitive(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + const PrimitivePtr &prim, const std::vector &inputs, + onnx::GraphProto *graph_proto); static onnx::TensorProto_DataType GetOnnxDataType(TypeId type_id); static onnx::TensorProto_DataType GetOutputType(const AnfNodePtr &node, int64_t output_index = -1); @@ -854,97 +883,102 @@ class OnnxExporter { mindspore::HashMap *op_merged_infos_ptr); void MatchAndMarkCNode(const FuncGraphPtr &func_graph, const CNodePtr &cnode, mindspore::HashMap *op_merged_infos_ptr); - void ExportNodes(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + void ExportNodes(const FuncGraphPtr &func_graph, std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportCNode(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + void ExportCNode(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimReshape(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimReduce(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimTranspose(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimStridedSlice(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); onnx::NodeProto *PrimResizeExportHelper(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto); void ExportPrimResizeNearestNeighbor(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimResizeBilinear(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimExpandDims(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimBatchMatMul(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimGeLU(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimGeLU(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimConcat(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimCast(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); - void ExportPrimPReLU(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); - void ExportPrimReLU6(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimCast(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimPReLU(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimReLU6(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimDepthwiseConv2d(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimTile(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimTile(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimSquare(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimGatherV2(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimTupleGetItem(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimTopK(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimTopK(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimBoundingBoxDecode(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimNMSWithMask(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimSplit(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimSplit(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimROIAlign(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimSlice(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimSlice(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimOnesLike(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimArgMaxWithValue(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimOneHot(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void PrimConv2DTransposeExportHelper(const CNodePtr &conv_node, const CNodePtr &bias_add_node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto); + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto); void ExportPrimConv2DTranspose(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimGreaterEqual(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimSqueeze(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportPrimLSTM(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimLSTM(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportPrimReverseV2(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportMergeConv(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); - void ExportMergeGemm(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimTensorCopySlices(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportPrimStack(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportMergeConv(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + void ExportMergeGemm(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportMergeBatchNorm(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportMergeMaxPoolWithArgmax(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportMergeLayerNorm(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); void ExportMergeConv2DTranspose(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::map *node_map_ptr, onnx::GraphProto *graph_proto); - void ExportOutput(const FuncGraphPtr &func_graph, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *graph_proto); - std::string GetNodeInputName(const AnfNodePtr &node, std::map *node_map_ptr, + void ExportOutput(const FuncGraphPtr &func_graph, const AnfNodePtr &return_arg, + std::map *node_map_ptr, onnx::GraphProto *graph_proto); + std::string GetNodeInputName(const AnfNodePtr &node, std::map *node_map_ptr, onnx::GraphProto *const graph_proto); void ConvertTupleToTensor(const ValuePtr &value, onnx::TensorProto *tensor_proto); @@ -953,7 +987,22 @@ class OnnxExporter { void AddOutputWithCast(onnx::NodeProto *node_proto, const std::string &output_name, onnx::TensorProto_DataType target_type, onnx::GraphProto *graph_proto); - size_t AllocateNodeIndex() { return ++onnx_node_index_; } + std::string GenerateUniqueName() { return std::to_string(++onnx_node_index_); } + std::string RegisterNodeWithUniqueName(const AnfNodePtr &node, std::map *node_map_ptr) { + auto name = GenerateUniqueName(); + (*node_map_ptr)[node] = name; + return name; + } + std::string GenerateUniqueParameterName(const ParameterPtr &node, std::map *node_map_ptr) { + auto node_name = node->ToString(); + MS_EXCEPTION_IF_CHECK_FAIL(node_name != "", "Cannot get the name of an ignored parameter"); + auto dup_iter = std::find_if(node_map_ptr->begin(), node_map_ptr->end(), + [&node_name](const auto &pair) { return pair.second == node_name; }); + if (dup_iter != node_map_ptr->end()) { + node_name = GenerateUniqueName() + node_name; + } + return node_name; + } void ResetNodeIndex() { onnx_node_index_ = 0; } @@ -966,6 +1015,8 @@ class OnnxExporter { onnx::ModelProto model_; size_t onnx_node_index_ = 0; + + std::map renamed_node_map_; }; std::string OnnxExporter::GetOnnxProtoString(const FuncGraphPtr &func_graph) { @@ -977,7 +1028,8 @@ std::string OnnxExporter::GetOnnxProtoString(const FuncGraphPtr &func_graph) { OpConvertRegistry::RegisterAllOpConverters(); InitModelInfo(); onnx::GraphProto *graph_proto = model_.mutable_graph(); - ExportFuncGraph(func_graph, graph_proto); + std::map node_map; + ExportFuncGraph(func_graph, &node_map, graph_proto); return model_.SerializeAsString(); } @@ -989,40 +1041,52 @@ void OnnxExporter::InitModelInfo() { opset_proto->set_version(ONNX_VERSION); } -void OnnxExporter::ExportFuncGraph(const FuncGraphPtr &func_graph, onnx::GraphProto *const graph_proto) { - std::map node_map; - +void OnnxExporter::ExportFuncGraph(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + onnx::GraphProto *const graph_proto, bool export_inputs) { MS_LOG(INFO) << "Begin exporting onnx model for graph " << func_graph->ToString(); - onnx_node_index_ = func_graph->parameters().size(); - // set graph name graph_proto->set_name(func_graph->ToString()); - // export parameters - // 1. all parameters (with or without default value) will be mapped to ONNX parameters - // 2. parameters with default value will mapped to ONNX initializers - ExportParameters(func_graph, graph_proto); + // export inputs if graph is not inlined + if (export_inputs) { + ExportInputs(func_graph, node_map_ptr, graph_proto); + } // export computational nodes and output nodes - ExportNodes(func_graph, &node_map, graph_proto); + ExportNodes(func_graph, node_map_ptr, graph_proto); MS_LOG(INFO) << "End exporting onnx model for graph " << func_graph->ToString(); } -void OnnxExporter::ExportParameters(const FuncGraphPtr &func_graph, onnx::GraphProto *const graph_proto) { +void OnnxExporter::ExportInputs(const FuncGraphPtr &func_graph, std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { for (auto ¶m : func_graph->parameters()) { const ParameterPtr param_ptr = dyn_cast(param); if (param_ptr == nullptr) { MS_LOG(EXCEPTION) << "Parameter '" << param->ToString() << "' could not cast to parameter."; } - // set onnx input. - if (!param_ptr->has_default()) { - onnx::ValueInfoProto *input_proto = graph_proto->add_input(); - input_proto->set_name(param_ptr->ToString()); - SetValueInfoType(param_ptr, input_proto); + if (param_ptr->has_default()) { + continue; } + + // set onnx input. + std::string name; + auto renamed_iter = renamed_node_map_.find(param_ptr); + if (renamed_iter != renamed_node_map_.end()) { + name = renamed_iter->second; + if (name == "") { + continue; + } + } else { + name = GenerateUniqueParameterName(param_ptr, node_map_ptr); + (*node_map_ptr)[param_ptr] = name; + } + + onnx::ValueInfoProto *input_proto = graph_proto->add_input(); + input_proto->set_name(name); + SetValueInfoType(param_ptr, input_proto); } } @@ -1086,7 +1150,7 @@ void OnnxExporter::MatchAndMark(const FuncGraphPtr &func_graph, const std::vecto auto &op_merged_infos = *op_merged_infos_ptr; for (auto &node : nodes) { - if (!node->isa() || IsIgnoredIdentityNode(node)) { + if (!node->isa() || IsZeroRefcountNode(node)) { continue; } auto cnode = node->cast(); @@ -1095,12 +1159,8 @@ void OnnxExporter::MatchAndMark(const FuncGraphPtr &func_graph, const std::vecto op_merged_infos[cnode].referred_count += 1; } for (auto &orig_input : cnode->inputs()) { - if (HasAbstractMonad(orig_input)) { - // Skip monad inputs. - continue; - } auto input = GetRealInput(orig_input); - if (!input->isa()) { + if (!input->isa() || IsZeroRefcountNode(input)) { continue; } // if the key `input` does not exist, just create a new one @@ -1163,31 +1223,17 @@ void OnnxExporter::MatchAndMarkCNode(const FuncGraphPtr &func_graph, const CNode * | +-- Parameter * | `-- ValueNode */ -void OnnxExporter::ExportNodes(const FuncGraphPtr &func_graph, std::map *node_map_ptr, +void OnnxExporter::ExportNodes(const FuncGraphPtr &func_graph, std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { std::vector nodes = TopoSort(func_graph->get_return(), SuccIncoming, AlwaysInclude); mindspore::HashMap op_merged_infos; MatchAndMark(func_graph, nodes, &op_merged_infos); - int count = -1; for (const AnfNodePtr &node : nodes) { - // skip when MakeTuple + UpdateState - count++; if (!node->isa()) { continue; } auto cnode = node->cast(); - if (cnode->IsApply(prim::kPrimMakeTuple)) { - size_t i = IntToSize(count + 1); - while (!nodes[i]->isa()) { - i++; - } - auto nextCNode = nodes[i]->cast(); - if (nextCNode->IsApply(prim::kPrimUpdateState) && - IsPrimitiveCNode(nextCNode->input(kTwoNum), prim::kPrimMakeTuple)) { - continue; - } - } auto iter = op_merged_infos.find(cnode); // the node is not referenced by any other nodes, skip it @@ -1200,7 +1246,7 @@ void OnnxExporter::ExportNodes(const FuncGraphPtr &func_graph, std::mapget_return()) { - ExportOutput(func_graph, cnode, node_map_ptr, graph_proto); + ExportOutput(func_graph, cnode->input(kOneNum), node_map_ptr, graph_proto); continue; } switch (merged_info.mode) { @@ -1230,15 +1276,14 @@ void OnnxExporter::ExportNodes(const FuncGraphPtr &func_graph, std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto name_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_shape = node->input(kTwoNum); std::string name_shape; if (input_shape->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[input_shape] = const_node_idx; + name_shape = RegisterNodeWithUniqueName(input_shape, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); - name_shape = std::to_string(const_node_idx); auto name = prim::kPrimReshape->name(); node_proto->set_name(name_shape + name); @@ -1253,23 +1298,22 @@ void OnnxExporter::ExportPrimReshape(const FuncGraphPtr &, const CNodePtr &node, MS_LOG(EXCEPTION) << "Need to insert op convert variable from tuple to tensor for Reshape."; } - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type(prim::kPrimReshape->name()); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(name_x); node_proto->add_input(name_shape); } void OnnxExporter::ExportPrimReduce(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto input_data = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_axis = node->input(kTwoNum); auto keep_dims = GetOpAttribute(node, "keep_dims"); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); std::string name; if (node->IsApply(prim::kPrimReduceSum)) { @@ -1300,22 +1344,21 @@ void OnnxExporter::ExportPrimReduce(const FuncGraphPtr &, const CNodePtr &node, MS_LOG(EXCEPTION) << "Need to insert op convert variable from tuple to attributes for " << name; } - AddReduceOp(name, input_data, std::to_string(node_idx), axes, keep_dims, graph_proto); + AddReduceOp(name, input_data, node_name, axes, keep_dims, graph_proto); } void OnnxExporter::ExportPrimTranspose(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_data = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_perm = node->input(kTwoNum); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); auto name = prim::kPrimTranspose->name(); - node_proto->set_name(std::to_string(node_idx) + name); + node_proto->set_name(node_name + name); node_proto->set_op_type(name); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(input_data); if (input_perm->isa()) { @@ -1339,105 +1382,93 @@ void OnnxExporter::ExportPrimTranspose(const FuncGraphPtr &, const CNodePtr &nod } } +/* + See: + - mindspore/ccsrc/backend/kernel_compiler/cpu/stridedslice_cpu_kernel.cc + - mindspore/ccsrc/backend/kernel_compiler/common_utils.cc + */ void OnnxExporter::ExportPrimStridedSlice(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_data = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); + auto name = node_name + prim::kPrimStridedSlice->name(); + auto begin = node->input(kTwoNum); - auto name = prim::kPrimStridedSlice->name(); - std::string name_begin; - if (begin->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[begin] = const_node_idx; - onnx::NodeProto *node_proto = graph_proto->add_node(); - name_begin = std::to_string(const_node_idx); - node_proto->add_output(name_begin); - - node_proto->set_op_type("Constant"); - onnx::AttributeProto *attr_proto = node_proto->add_attribute(); - attr_proto->set_name("value"); - - attr_proto->set_type(onnx::AttributeProto_AttributeType_TENSOR); - ConvertTupleToTensor(dyn_cast(begin)->value(), attr_proto->mutable_t()); - } else { + if (!begin->isa()) { MS_LOG(EXCEPTION) << "The input begin of StridedSlice is not a ValueNode! " << "Need to insert op convert variable from tuple to tensor for " << name; } + auto begin_value_node = dyn_cast(begin); + auto begin_value = GetValue>(begin_value_node->value()); + auto begin_ignore_mask = GetOpAttribute(node, "begin_mask"); + for (size_t i = 0; i < begin_value.size(); ++i) { + if (begin_ignore_mask & (1 << i)) { + begin_value[i] = 0; + } + } auto end = node->input(kThreeNum); - std::string name_end; - if (end->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[end] = const_node_idx; - onnx::NodeProto *node_proto = graph_proto->add_node(); - name_end = std::to_string(const_node_idx); - node_proto->add_output(name_end); - - node_proto->set_op_type("Constant"); - onnx::AttributeProto *attr_proto = node_proto->add_attribute(); - attr_proto->set_name("value"); - - attr_proto->set_type(onnx::AttributeProto_AttributeType_TENSOR); - ConvertTupleToTensor(dyn_cast(end)->value(), attr_proto->mutable_t()); - } else { + if (!end->isa()) { MS_LOG(EXCEPTION) << "The input end of StridedSlice is not a ValueNode! " << "Need to insert op convert variable from tuple to tensor for " << name; } - - auto x_shape = dyn_cast(node->input(1)->Shape()); - int size = SizeToInt(x_shape->shape().size()); - std::vector axes_value; - ValuePtr axes_value_ptr = nullptr; - for (int i = 0; i < size; ++i) { - axes_value.push_back(i); + auto end_value_node = dyn_cast(end); + auto end_value = GetValue>(end_value_node->value()); + const auto &x_shape = dyn_cast(node->input(kOneNum)->Shape())->shape(); + auto end_ignore_mask = GetOpAttribute(node, "end_mask"); + for (size_t i = 0; i < end_value.size(); ++i) { + if (end_ignore_mask & (1 << i)) { + end_value[i] = x_shape[i]; + } + } + + std::vector axes_value; + for (size_t i = 0; i < x_shape.size(); ++i) { + axes_value.push_back(static_cast(i)); } - axes_value_ptr = MakeValue>(axes_value); - auto axes = NewValueNode(axes_value_ptr)->cast(); - std::string name_axes; - auto const_node_idx_axes = AllocateNodeIndex(); - (*node_map_ptr)[axes] = const_node_idx_axes; - onnx::NodeProto *node_proto_axes = graph_proto->add_node(); - name_axes = std::to_string(const_node_idx_axes); - node_proto_axes->add_output(name_axes); - node_proto_axes->set_op_type("Constant"); - onnx::AttributeProto *attr_proto_axes = node_proto_axes->add_attribute(); - attr_proto_axes->set_name("value"); - attr_proto_axes->set_type(onnx::AttributeProto_AttributeType_TENSOR); - ConvertTupleToTensor(dyn_cast(axes)->value(), attr_proto_axes->mutable_t()); auto strides = node->input(kFourNum); - std::string name_strides; - if (strides->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[strides] = const_node_idx; - onnx::NodeProto *node_proto = graph_proto->add_node(); - name_strides = std::to_string(const_node_idx); - node_proto->add_output(name_strides); - - node_proto->set_op_type("Constant"); - onnx::AttributeProto *attr_proto_steps = node_proto->add_attribute(); - attr_proto_steps->set_name("value"); - attr_proto_steps->set_type(onnx::AttributeProto_AttributeType_TENSOR); - ConvertTupleToTensor(dyn_cast(strides)->value(), attr_proto_steps->mutable_t()); - } else { + if (!strides->isa()) { MS_LOG(EXCEPTION) << "The input strides of StridedSlice is not a ValueNode! " << "Need to insert op convert variable from tuple to tensor for " << name; } + auto strides_value_node = dyn_cast(strides); + auto strides_value = GetValue>(strides_value_node->value()); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; - onnx::NodeProto *node_proto = graph_proto->add_node(); - node_proto->set_op_type("Slice"); - node_proto->add_output(std::to_string(node_idx)); - node_proto->add_input(input_data); - node_proto->add_input(name_begin); - node_proto->add_input(name_end); - node_proto->add_input(name_axes); - node_proto->add_input(name_strides); + auto shrink_axis_mask = GetOpAttribute(node, "shrink_axis_mask"); + for (size_t i = 0; i < end_value.size(); ++i) { + if (shrink_axis_mask & (1 << i)) { + strides_value[i] = end_value[i] > begin_value[i] ? 1 : -1; + end_value[i] = begin_value[i] + strides_value[i]; + } + } + + auto slice_name = node_name; + if (shrink_axis_mask != 0) { + slice_name = node_name + "__reshape"; + } + + AddSliceOp(input_data, slice_name, begin_value, end_value, axes_value, strides_value, graph_proto); + + if (shrink_axis_mask != 0) { + onnx::NodeProto *squeeze_op = graph_proto->add_node(); + squeeze_op->set_op_type("Squeeze"); + squeeze_op->add_input(slice_name); + squeeze_op->add_output(node_name); + onnx::AttributeProto *axes_attr = squeeze_op->add_attribute(); + axes_attr->set_name("axes"); + axes_attr->set_type(onnx::AttributeProto_AttributeType_INTS); + for (size_t i = 0; i < x_shape.size(); ++i) { + if (shrink_axis_mask & (1 << i)) { + axes_attr->add_ints(i); + } + } + } } onnx::NodeProto *OnnxExporter::PrimResizeExportHelper(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_data = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto x_shape = dyn_cast(node->input(kOneNum)->Shape()); @@ -1462,12 +1493,9 @@ onnx::NodeProto *OnnxExporter::PrimResizeExportHelper(const FuncGraphPtr &, cons } auto resize_size_ptr = MakeValue>(resize_size); auto size = NewValueNode(resize_size_ptr)->cast(); - std::string name_size; - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[size] = const_node_idx; + auto name_size = RegisterNodeWithUniqueName(size, node_map_ptr); onnx::NodeProto *node_proto_size = graph_proto->add_node(); - name_size = std::to_string(const_node_idx); node_proto_size->add_output(name_size); node_proto_size->set_op_type("Constant"); onnx::AttributeProto *attr_proto = node_proto_size->add_attribute(); @@ -1475,25 +1503,24 @@ onnx::NodeProto *OnnxExporter::PrimResizeExportHelper(const FuncGraphPtr &, cons attr_proto->set_type(onnx::AttributeProto_AttributeType_TENSOR); ConvertTupleToTensor(resize_size_ptr, attr_proto->mutable_t()); - auto node_idx = AllocateNodeIndex(); + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::TensorProto *roi_initializer_proto = graph_proto->add_initializer(); - auto roi_name = std::to_string(node_idx) + "roi_initializer"; + auto roi_name = node_name + "roi_initializer"; roi_initializer_proto->set_name(roi_name); roi_initializer_proto->set_data_type(GetOnnxDataType(kNumberTypeFloat32)); roi_initializer_proto->add_dims(0); onnx::TensorProto *scales_initializer_proto = graph_proto->add_initializer(); - auto scales_name = std::to_string(node_idx) + "scales_initializer"; + auto scales_name = node_name + "scales_initializer"; scales_initializer_proto->set_name(scales_name); scales_initializer_proto->set_data_type(GetOnnxDataType(kNumberTypeFloat32)); scales_initializer_proto->add_dims(0); - (*node_map_ptr)[node] = node_idx; onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Resize"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(input_data); node_proto->add_input(roi_name); node_proto->add_input(scales_name); @@ -1503,7 +1530,7 @@ onnx::NodeProto *OnnxExporter::PrimResizeExportHelper(const FuncGraphPtr &, cons } void OnnxExporter::ExportPrimResizeNearestNeighbor(const FuncGraphPtr &graph, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { onnx::NodeProto *node_proto = PrimResizeExportHelper(graph, node, node_map_ptr, graph_proto); @@ -1525,7 +1552,7 @@ void OnnxExporter::ExportPrimResizeNearestNeighbor(const FuncGraphPtr &graph, co } void OnnxExporter::ExportPrimResizeBilinear(const FuncGraphPtr &graph, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { onnx::NodeProto *node_proto = PrimResizeExportHelper(graph, node, node_map_ptr, graph_proto); @@ -1545,7 +1572,7 @@ void OnnxExporter::ExportPrimResizeBilinear(const FuncGraphPtr &graph, const CNo // MindSpore ExpandDims -> ONNX Reshape void OnnxExporter::ExportPrimExpandDims(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto axis = GetInt64Value(node->input(kTwoNum)); @@ -1565,10 +1592,8 @@ void OnnxExporter::ExportPrimExpandDims(const FuncGraphPtr &, const CNodePtr &no std::string name_shape; if (shape->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[shape] = const_node_idx; + name_shape = RegisterNodeWithUniqueName(shape, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); - name_shape = std::to_string(const_node_idx); node_proto->add_output(name_shape); node_proto->set_op_type("Constant"); onnx::AttributeProto *attr_proto = node_proto->add_attribute(); @@ -1580,18 +1605,17 @@ void OnnxExporter::ExportPrimExpandDims(const FuncGraphPtr &, const CNodePtr &no MS_LOG(EXCEPTION) << "Need to insert op convert variable from tuple to tensor for " << name; } - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Reshape"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(input_x); node_proto->add_input(name_shape); } // MindSpore BatchMatMul -> ONNX Transpose + MatMul void OnnxExporter::ExportPrimBatchMatMul(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_y = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); @@ -1607,10 +1631,10 @@ void OnnxExporter::ExportPrimBatchMatMul(const FuncGraphPtr &, const CNodePtr &n if (transpose_a) { auto input_x_shape = dyn_cast(node->input(kOneNum)->Shape()); // Add Transpose node after input_x of BatchMatMul - auto transpose_input_x_index = AllocateNodeIndex(); + transpose_input_x_name = GenerateUniqueName(); onnx::NodeProto *transpose_inputx_node_proto = graph_proto->add_node(); transpose_inputx_node_proto->add_input(input_x); - transpose_inputx_node_proto->add_output(std::to_string(transpose_input_x_index)); + transpose_inputx_node_proto->add_output(transpose_input_x_name); transpose_inputx_node_proto->set_op_type(prim::kPrimTranspose->name()); onnx::AttributeProto *attr_proto = transpose_inputx_node_proto->add_attribute(); attr_proto->set_name("perm"); @@ -1620,15 +1644,14 @@ void OnnxExporter::ExportPrimBatchMatMul(const FuncGraphPtr &, const CNodePtr &n } attr_proto->add_ints(SizeToLong(input_x_shape->shape().size()) - IntToLong(kOneNum)); attr_proto->add_ints(SizeToLong(input_x_shape->shape().size()) - IntToLong(kTwoNum)); - transpose_input_x_name = std::to_string(transpose_input_x_index); } if (transpose_b) { auto input_y_shape = dyn_cast(node->input(kTwoNum)->Shape()); // Add Transpose node after input_y of BatchMatMul - auto transpose_input_y_index = AllocateNodeIndex(); + transpose_input_y_name = GenerateUniqueName(); onnx::NodeProto *transpose_inputy_node_proto = graph_proto->add_node(); transpose_inputy_node_proto->add_input(input_y); - transpose_inputy_node_proto->add_output(std::to_string(transpose_input_y_index)); + transpose_inputy_node_proto->add_output(transpose_input_y_name); transpose_inputy_node_proto->set_op_type(prim::kPrimTranspose->name()); onnx::AttributeProto *attr_proto = transpose_inputy_node_proto->add_attribute(); attr_proto->set_name("perm"); @@ -1638,15 +1661,13 @@ void OnnxExporter::ExportPrimBatchMatMul(const FuncGraphPtr &, const CNodePtr &n } attr_proto->add_ints(SizeToLong(input_y_shape->shape().size()) - IntToLong(kOneNum)); attr_proto->add_ints(SizeToLong(input_y_shape->shape().size()) - IntToLong(kTwoNum)); - transpose_input_y_name = std::to_string(transpose_input_y_index); } - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("MatMul"); - node_proto->add_output(std::to_string(node_idx)); - node_proto->set_name(std::to_string(node_idx) + "MatMul"); + node_proto->add_output(node_name); + node_proto->set_name(node_name + "MatMul"); if (transpose_a) { node_proto->add_input(transpose_input_x_name); } else { @@ -1661,60 +1682,58 @@ void OnnxExporter::ExportPrimBatchMatMul(const FuncGraphPtr &, const CNodePtr &n // MindSpore GeLU -> ONNX 0.5 * X * (1.0 + tanh((sqrt(2/pi) * (x + 0.044715 * pow(x, 3))))) void OnnxExporter::ExportPrimGeLU(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto input_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto onnx_type = GetOutputType(node->input(kOneNum)); // Add pow node - auto pow_name = std::to_string(AllocateNodeIndex()); + auto pow_name = GenerateUniqueName(); auto exp_node_name = pow_name + "exponent_initializer"; AddFloatTensor1DInitializer(exp_node_name, {3.0}, onnx_type, graph_proto); AddOp("Pow", {input_x, exp_node_name}, {pow_name}, graph_proto); // Add first Mul Node - auto fmul_name = std::to_string(AllocateNodeIndex()); + auto fmul_name = GenerateUniqueName(); auto fmul_input_node_name = fmul_name + "input_y_for_mul_initializer"; AddFloatTensor1DInitializer(fmul_input_node_name, {0.044715}, onnx_type, graph_proto); AddOp("Mul", {pow_name, fmul_input_node_name}, {fmul_name}, graph_proto); // Add first Add node - auto fadd_name = std::to_string(AllocateNodeIndex()); + auto fadd_name = GenerateUniqueName(); AddOp("Add", {input_x, fmul_name}, {fadd_name}, graph_proto); // Add second Mul Node - auto smul_name = std::to_string(AllocateNodeIndex()); + auto smul_name = GenerateUniqueName(); auto smul_input_node_name = smul_name + "input_y_for_smul_initializer"; AddFloatTensor1DInitializer(smul_input_node_name, {0.7978845608}, onnx_type, graph_proto); AddOp("Mul", {fadd_name, smul_input_node_name}, {smul_name}, graph_proto); // Add tanh node - auto tanh_name = std::to_string(AllocateNodeIndex()); + auto tanh_name = GenerateUniqueName(); AddOp("Tanh", {smul_name}, {tanh_name}, graph_proto); // Add second Add node - auto sadd_name = std::to_string(AllocateNodeIndex()); + auto sadd_name = GenerateUniqueName(); auto sadd_input_node_name = sadd_name + "input_y_for_sadd_initializer"; AddFloatTensor1DInitializer(sadd_input_node_name, {1.0}, onnx_type, graph_proto); AddOp("Add", {tanh_name, sadd_input_node_name}, {sadd_name}, graph_proto); // Add third Mul Node - auto tmul_name = std::to_string(AllocateNodeIndex()); + auto tmul_name = GenerateUniqueName(); auto tmul_input_node_name = tmul_name + "input_y_for_tmul_initializer"; AddFloatTensor1DInitializer(tmul_input_node_name, {0.5}, onnx_type, graph_proto); AddOp("Mul", {sadd_name, tmul_input_node_name}, {tmul_name}, graph_proto); // Add fourth Mul Node - auto fomul_node_idx = AllocateNodeIndex(); - auto fomul_node_name = std::to_string(fomul_node_idx); + auto fomul_node_name = RegisterNodeWithUniqueName(node, node_map_ptr); AddOp("Mul", {input_x, tmul_name}, {fomul_node_name}, graph_proto); - - (*node_map_ptr)[node] = fomul_node_idx; } void OnnxExporter::ExportPrimConcat(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); // Get inputs first: otherwise if an input is a constant, topological order will break auto input_node = node->input(kOneNum)->cast(); @@ -1729,19 +1748,19 @@ void OnnxExporter::ExportPrimConcat(const FuncGraphPtr &, const CNodePtr &node, input_names.push_back(input_data); } - AddConcatOp(input_names, std::to_string(node_idx), GetOpAttribute(node, "axis"), graph_proto); + AddConcatOp(input_names, node_name, GetOpAttribute(node, "axis"), graph_proto); } void OnnxExporter::ExportPrimCast(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto input_data = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_type = node->input(kTwoNum); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type(prim::kPrimCast->name()); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(input_data); if (input_type->isa()) { @@ -1758,7 +1777,8 @@ void OnnxExporter::ExportPrimCast(const FuncGraphPtr &, const CNodePtr &node, } void OnnxExporter::ExportPrimPReLU(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto input_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_slope = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); @@ -1769,10 +1789,10 @@ void OnnxExporter::ExportPrimPReLU(const FuncGraphPtr &, const CNodePtr &node, // format of x is NCHW, input format is NCHW, if length of input_slope is 1, insert Unsqueeze [1,2] if (x_shape->shape().size() == kFourNum && slope_shape->shape().size() == kOneNum) { - auto node_idx = AllocateNodeIndex(); + auto node_name = GenerateUniqueName(); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Unsqueeze"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); onnx::AttributeProto *attr_proto = node_proto->add_attribute(); attr_proto->set_type(onnx::AttributeProto_AttributeType_INTS); @@ -1781,23 +1801,21 @@ void OnnxExporter::ExportPrimPReLU(const FuncGraphPtr &, const CNodePtr &node, attr_proto->add_ints(kTwoNum); node_proto->add_input(input_slope); - input_slope = std::to_string(node_idx); + input_slope = node_name; } - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("PRelu"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(input_x); node_proto->add_input(input_slope); } void OnnxExporter::ExportPrimReLU6(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_x_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto onnx_input_type = GetOutputType(node->input(kOneNum)); @@ -1805,7 +1823,7 @@ void OnnxExporter::ExportPrimReLU6(const FuncGraphPtr &, const CNodePtr &node, } void OnnxExporter::ExportPrimDepthwiseConv2d(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto input_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_w = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); @@ -1820,9 +1838,9 @@ void OnnxExporter::ExportPrimDepthwiseConv2d(const FuncGraphPtr &, const CNodePt MS_LOG(EXCEPTION) << "DepthwiseConv2d weight shape[0] != 1 and shape[1] != 1, cannot reshape"; } // create w_shape constant node - auto node_idx = AllocateNodeIndex(); + auto node_name = GenerateUniqueName(); onnx::NodeProto *node_proto = graph_proto->add_node(); - std::string name_w_shape = std::to_string(node_idx); + auto name_w_shape = node_name; node_proto->add_output(name_w_shape); node_proto->set_op_type("Constant"); // create Value Tensor @@ -1839,22 +1857,21 @@ void OnnxExporter::ExportPrimDepthwiseConv2d(const FuncGraphPtr &, const CNodePt tensor_proto->add_int64_data(w_shape->shape()[kThreeNum]); // add reshape node - node_idx = AllocateNodeIndex(); + node_name = GenerateUniqueName(); node_proto = graph_proto->add_node(); node_proto->set_op_type(prim::kPrimReshape->name()); node_proto->add_input(input_w); node_proto->add_input(name_w_shape); - input_w = std::to_string(node_idx); + input_w = node_name; node_proto->add_output(input_w); // add conv node - node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + node_name = RegisterNodeWithUniqueName(node, node_map_ptr); node_proto = graph_proto->add_node(); node_proto->set_op_type("Conv"); node_proto->add_input(input_x); node_proto->add_input(input_w); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); // set attributes AnfNodePtr op = node->input(0); auto op_value = dyn_cast(op); @@ -1896,15 +1913,14 @@ void OnnxExporter::ExportPrimDepthwiseConv2d(const FuncGraphPtr &, const CNodePt } void OnnxExporter::ExportPrimTile(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto name_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto multiples = node->input(kTwoNum); std::string name_multiples; if (multiples->isa()) { - auto const_node_idx = AllocateNodeIndex(); - (*node_map_ptr)[multiples] = const_node_idx; onnx::NodeProto *node_proto = graph_proto->add_node(); - name_multiples = std::to_string(const_node_idx); + name_multiples = RegisterNodeWithUniqueName(multiples, node_map_ptr); node_proto->add_output(name_multiples); node_proto->set_op_type("Constant"); onnx::AttributeProto *attr_proto = node_proto->add_attribute(); @@ -1916,22 +1932,20 @@ void OnnxExporter::ExportPrimTile(const FuncGraphPtr &, const CNodePtr &node, MS_LOG(EXCEPTION) << "Need to insert op convert variable from tuple to tensor for Tile."; } - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Tile"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(name_x); node_proto->add_input(name_multiples); } void OnnxExporter::ExportPrimSquare(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto name_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); - std::string name_exponent; - auto const_node_idx = AllocateNodeIndex(); + auto name_exponent = GenerateUniqueName(); onnx::NodeProto *node_proto_exp = graph_proto->add_node(); - name_exponent = std::to_string(const_node_idx); node_proto_exp->add_output(name_exponent); node_proto_exp->set_op_type("Constant"); @@ -1945,25 +1959,24 @@ void OnnxExporter::ExportPrimSquare(const FuncGraphPtr &, const CNodePtr &node, tensor_proto->set_data_type(GetOnnxDataType(kNumberTypeFloat32)); tensor_proto->add_float_data(exponent_value); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Pow"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(name_x); node_proto->add_input(name_exponent); } void OnnxExporter::ExportPrimGatherV2(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto name_x = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto name_indices = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); auto axis = node->input(kThreeNum)->cast()->value(); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Gather"); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); node_proto->add_input(name_x); node_proto->add_input(name_indices); onnx::AttributeProto *attr_proto = node_proto->add_attribute(); @@ -1983,29 +1996,27 @@ void OnnxExporter::ExportPrimGatherV2(const FuncGraphPtr &, const CNodePtr &node See OnnxExporter::ExportPrimTopK for a usage example */ void OnnxExporter::ExportPrimTupleGetItem(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto index = GetInt64Value(node->input(kTwoNum)); auto input_node_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_name = MakeOutputName(input_node_name, index); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Identity"); node_proto->add_input(input_name); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); } void OnnxExporter::ExportPrimTopK(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto x_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto k_input_name = node_name + "k_initializer"; auto k = GetInt64Value(node->input(kTwoNum)); @@ -2030,11 +2041,9 @@ void OnnxExporter::ExportPrimTopK(const FuncGraphPtr &, const CNodePtr &node, // Based on mindspore/ccsrc/backend/kernel_compiler/cpu/boundingbox_decode_cpu_kernel.cc void OnnxExporter::ExportPrimBoundingBoxDecode(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto anchor_bbox_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto deltas_input_name = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); @@ -2109,11 +2118,9 @@ void OnnxExporter::ExportPrimBoundingBoxDecode(const FuncGraphPtr &, const CNode } void OnnxExporter::ExportPrimNMSWithMask(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto bboxes_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto iou_threshold = GetOpAttribute(node, "iou_threshold"); @@ -2212,10 +2219,9 @@ void OnnxExporter::ExportPrimNMSWithMask(const FuncGraphPtr &, const CNodePtr &n } void OnnxExporter::ExportPrimSplit(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto axis = GetOpAttribute(node, "axis"); @@ -2259,10 +2265,9 @@ void OnnxExporter::ExportPrimSplit(const FuncGraphPtr &, const CNodePtr &node, * MS has two ROI end modes, implemented with pre-processing */ void OnnxExporter::ExportPrimROIAlign(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto features_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto rois_input_name = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); auto onnx_input_type = GetOutputType(node->input(kOneNum)); @@ -2323,10 +2328,9 @@ void OnnxExporter::ExportPrimROIAlign(const FuncGraphPtr &, const CNodePtr &node } void OnnxExporter::ExportPrimSlice(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_x_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto begin_input_name = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); auto size_input_name = GetNodeInputName(node->input(kThreeNum), node_map_ptr, graph_proto); @@ -2337,10 +2341,9 @@ void OnnxExporter::ExportPrimSlice(const FuncGraphPtr &, const CNodePtr &node, } void OnnxExporter::ExportPrimOnesLike(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_x_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto shape_name = node_name + "shape"; @@ -2374,11 +2377,9 @@ void OnnxExporter::ExportPrimOnesLike(const FuncGraphPtr &, const CNodePtr &node } void OnnxExporter::ExportPrimArgMaxWithValue(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_x_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto axis = GetOpAttribute(node, "axis"); auto keep_dims = GetOpAttribute(node, "keep_dims"); @@ -2406,10 +2407,9 @@ void OnnxExporter::ExportPrimArgMaxWithValue(const FuncGraphPtr &, const CNodePt } void OnnxExporter::ExportPrimOneHot(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto indices_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto depth_input_name = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); auto on_input_name = GetNodeInputName(node->input(kThreeNum), node_map_ptr, graph_proto); @@ -2448,16 +2448,16 @@ void OnnxExporter::ExportPrimOneHot(const FuncGraphPtr &, const CNodePtr &node, it is not possible to change the output shape in runtime */ void OnnxExporter::PrimConv2DTransposeExportHelper(const CNodePtr &conv_node, const CNodePtr &bias_add_node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); + std::string node_name; std::vector inputs{conv_node->input(kOneNum), conv_node->input(kTwoNum)}; if (bias_add_node != nullptr) { inputs.push_back(bias_add_node->input(kTwoNum)); - (*node_map_ptr)[bias_add_node] = node_idx; + node_name = RegisterNodeWithUniqueName(bias_add_node, node_map_ptr); } else { - (*node_map_ptr)[conv_node] = node_idx; + node_name = RegisterNodeWithUniqueName(conv_node, node_map_ptr); } onnx::NodeProto *node_proto = graph_proto->add_node(); @@ -2465,7 +2465,7 @@ void OnnxExporter::PrimConv2DTransposeExportHelper(const CNodePtr &conv_node, co for (const auto &input : inputs) { node_proto->add_input(GetNodeInputName(input, node_map_ptr, graph_proto)); } - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); auto prim = GetPrimitive(conv_node); auto attrs_convert_info = @@ -2504,17 +2504,15 @@ void OnnxExporter::PrimConv2DTransposeExportHelper(const CNodePtr &conv_node, co } void OnnxExporter::ExportPrimConv2DTranspose(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *graph_proto) { PrimConv2DTransposeExportHelper(node, nullptr, node_map_ptr, graph_proto); } void OnnxExporter::ExportPrimGreaterEqual(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_x_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); auto input_y_name = GetNodeInputName(node->input(kTwoNum), node_map_ptr, graph_proto); @@ -2525,16 +2523,16 @@ void OnnxExporter::ExportPrimGreaterEqual(const FuncGraphPtr &, const CNodePtr & } void OnnxExporter::ExportPrimSqueeze(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); onnx::NodeProto *node_proto = graph_proto->add_node(); node_proto->set_op_type("Squeeze"); node_proto->add_input(input_name); - node_proto->add_output(std::to_string(node_idx)); + node_proto->add_output(node_name); auto axes = GetOpAttributePtr(node, "axis"); auto axes_value = GetValue>(axes); @@ -2639,10 +2637,9 @@ void ExportLSTMWeights(const CNodePtr &node, const std::string &node_name, const } void OnnxExporter::ExportPrimLSTM(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); - (*node_map_ptr)[node] = node_idx; + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); // MS inputs auto x_input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); @@ -2715,12 +2712,10 @@ void OnnxExporter::ExportPrimLSTM(const FuncGraphPtr &func_graph, const CNodePtr } void OnnxExporter::ExportPrimReverseV2(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto output = RegisterNodeWithUniqueName(node, node_map_ptr); auto input = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); - auto output = std::to_string(node_idx); auto axes_ptr = GetOpAttributePtr(node, "axis"); auto axes_vec = GetValue>(axes_ptr); @@ -2736,10 +2731,105 @@ void OnnxExporter::ExportPrimReverseV2(const FuncGraphPtr &, const CNodePtr &nod AddSliceOp(input, output, starts_vec, ends_vec, axes_vec, steps_vec, graph_proto); } +void OnnxExporter::ExportPrimTensorCopySlices(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, + onnx::GraphProto *graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); + + auto x_input = node->input(kOneNum); + auto value_input = node->input(kTwoNum); + + auto x_input_name = GetNodeInputName(x_input, node_map_ptr, graph_proto); + auto value_input_name = GetNodeInputName(value_input, node_map_ptr, graph_proto); + + const auto &x_shape = dyn_cast(x_input->Shape())->shape(); + const auto &value_shape = dyn_cast(value_input->Shape())->shape(); + + auto begin_node = dyn_cast(node->input(kThreeNum)); + MS_EXCEPTION_IF_NULL(begin_node); + auto begin = GetValue>(begin_node->value()); + + auto end_node = dyn_cast(node->input(kFourNum)); + MS_EXCEPTION_IF_NULL(end_node); + auto end = GetValue>(end_node->value()); + + auto strides_node = dyn_cast(node->input(kFiveNum)); + MS_EXCEPTION_IF_NULL(strides_node); + auto strides = GetValue>(strides_node->value()); + + MS_EXCEPTION_IF_CHECK_FAIL( + begin.size() == end.size() && end.size() == strides.size() && strides.size() <= x_shape.size(), + "Sizes of begin, end, and strides must be equal"); + // MindSpore only allows contuguous slices of memory + // Contiguous slice size follows the pattern: [1, ..., 1, n, :, ..., :] + bool found_slice = false; + for (size_t i = 0; i < begin.size(); ++i) { + int64_t dim = end[i] - begin[i]; + if (!found_slice && dim != 1) { + found_slice = true; + } else if (found_slice && dim != x_shape[i]) { + MS_LOG(EXCEPTION) << "Slice must be contiguous"; + } + } + for (auto stride : strides) { + MS_EXCEPTION_IF_CHECK_FAIL(stride == 1, "Slice must be contiguous"); + } + + size_t flat_begin_index = RavelIndex(begin, x_shape); + + std::vector end_inclusive; + std::transform(end.begin(), end.end(), std::back_inserter(end_inclusive), [](auto x) { return x - 1; }); + std::transform(x_shape.begin() + end.size(), x_shape.end(), std::back_inserter(end_inclusive), + [](auto x) { return x - 1; }); + size_t flat_end_index = RavelIndex(end_inclusive, x_shape) + 1; + + size_t x_size = std::accumulate(x_shape.begin(), x_shape.end(), 1, std::multiplies()); + size_t value_size = std::accumulate(value_shape.begin(), value_shape.end(), 1, std::multiplies()); + MS_EXCEPTION_IF_CHECK_FAIL(value_size == flat_end_index - flat_begin_index, "Cannot copy 'value' to target slice"); + + auto flat_x_name = node_name + "_flat_x"; + AddReshapeOp(x_input_name, flat_x_name, {-1}, graph_proto); + auto begin_slice_name = node_name + "_begin_slice"; + AddSliceOp(flat_x_name, begin_slice_name, {0}, {static_cast(flat_begin_index)}, {0}, {1}, graph_proto); + auto end_slice_name = node_name + "_end_slice"; + AddSliceOp(flat_x_name, end_slice_name, {static_cast(flat_end_index)}, {static_cast(x_size)}, {0}, + {1}, graph_proto); + + auto flat_value_name = node_name + "_flat_value"; + AddReshapeOp(value_input_name, flat_value_name, {-1}, graph_proto); + + auto flat_result_name = node_name + "_flat_result"; + AddConcatOp({begin_slice_name, flat_value_name, end_slice_name}, flat_result_name, 0, graph_proto); + AddReshapeOp(flat_result_name, node_name, x_shape, graph_proto); +} + +void OnnxExporter::ExportPrimStack(const FuncGraphPtr &func_graph, const CNodePtr &node, + std::map *node_map_ptr, onnx::GraphProto *graph_proto) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); + + auto input_name = GetNodeInputName(node->input(kOneNum), node_map_ptr, graph_proto); + + onnx::NodeProto *node_proto = graph_proto->add_node(); + node_proto->set_name(node_name + "Stack"); + node_proto->set_op_type("ConcatFromSequence"); + node_proto->add_input(input_name); + node_proto->add_output(node_name); + + onnx::AttributeProto *axis_proto = node_proto->add_attribute(); + axis_proto->set_name("axis"); + axis_proto->set_type(onnx::AttributeProto_AttributeType_INT); + axis_proto->set_i(GetOpAttribute(node, "axis")); + + onnx::AttributeProto *new_axis_proto = node_proto->add_attribute(); + new_axis_proto->set_name("new_axis"); + new_axis_proto->set_type(onnx::AttributeProto_AttributeType_INT); + new_axis_proto->set_i(true); +} + void OnnxExporter::ExportCNode(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { using ExportFunc = std::function *, onnx::GraphProto *const)>; + std::map *, onnx::GraphProto *const)>; static std::vector> export_table = { {prim::kPrimReshape, &OnnxExporter::ExportPrimReshape}, {prim::kPrimReduceMean, &OnnxExporter::ExportPrimReduce}, @@ -2774,6 +2864,8 @@ void OnnxExporter::ExportCNode(const FuncGraphPtr &func_graph, const CNodePtr &n {prim::kPrimGeLU, &OnnxExporter::ExportPrimGeLU}, {prim::kPrimLstm, &OnnxExporter::ExportPrimLSTM}, {prim::kPrimReverseV2, &OnnxExporter::ExportPrimReverseV2}, + {prim::kPrimTensorCopySlices, &OnnxExporter::ExportPrimTensorCopySlices}, + {prim::kPrimStack, &OnnxExporter::ExportPrimStack}, }; auto iter = std::find_if(export_table.begin(), export_table.end(), @@ -2862,9 +2954,9 @@ void OnnxExporter::AddOutputWithCast(onnx::NodeProto *node_proto, const std::str } } -size_t OnnxExporter::ExportPrimitive(const FuncGraphPtr &, std::map *node_map_ptr, - const PrimitivePtr &prim, const std::vector &inputs, - onnx::GraphProto *const graph_proto) { +std::string OnnxExporter::ExportPrimitive(const FuncGraphPtr &, std::map *node_map_ptr, + const PrimitivePtr &prim, const std::vector &inputs, + onnx::GraphProto *const graph_proto) { auto op_map = OpConvertRegistry::GetOpConvertMap(); MS_EXCEPTION_IF_NULL(prim); auto op_iter = op_map.find(prim->name()); @@ -2880,8 +2972,7 @@ size_t OnnxExporter::ExportPrimitive(const FuncGraphPtr &, std::mapsecond; - auto node_idx = AllocateNodeIndex(); - auto node_name = std::to_string(node_idx); + auto node_name = GenerateUniqueName(); std::vector output_cast_types(op_convert_info.num_outputs(), onnx::TensorProto_DataType_UNDEFINED); @@ -2944,11 +3035,12 @@ size_t OnnxExporter::ExportPrimitive(const FuncGraphPtr &, std::mapset_name(attr.onnx_attr_name()); attr.fn_gen_attr()(attr_value, attr.onnx_attr_type(), onnx_attr_proto, prim); } - return node_idx; + return node_name; } void OnnxExporter::ExportMergeConv(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto conv_node = dyn_cast(node->input(kOneNum)); auto input_x = conv_node->input(kOneNum); // conv input x auto input_w = conv_node->input(kTwoNum); // conv weight(filter) @@ -2960,7 +3052,8 @@ void OnnxExporter::ExportMergeConv(const FuncGraphPtr &func_graph, const CNodePt } void OnnxExporter::ExportMergeGemm(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + std::map *node_map_ptr, + onnx::GraphProto *const graph_proto) { auto matmul_node = dyn_cast(node->input(kOneNum)); auto input_x = matmul_node->input(kOneNum); // matmul input x auto input_y = matmul_node->input(kTwoNum); // matmul input y @@ -2972,7 +3065,7 @@ void OnnxExporter::ExportMergeGemm(const FuncGraphPtr &func_graph, const CNodePt } void OnnxExporter::ExportMergeBatchNorm(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto batch_norm_node = dyn_cast(node->input(kOneNum)); @@ -2984,9 +3077,7 @@ void OnnxExporter::ExportMergeBatchNorm(const FuncGraphPtr &func_graph, const CN auto onnx_type = GetOutputType(batch_norm_node->input(kOneNum)); - auto output_index = AllocateNodeIndex(); - auto output_name = std::to_string(output_index); - (*node_map_ptr)[node] = output_index; + auto output_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto input_shape_ptr = batch_norm_node->input(kOneNum)->Shape(); auto input_shape = input_shape_ptr->cast()->shape(); @@ -3017,7 +3108,7 @@ void OnnxExporter::ExportMergeBatchNorm(const FuncGraphPtr &func_graph, const CN } void OnnxExporter::ExportMergeMaxPoolWithArgmax(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto maxpool_with_argmax_node = dyn_cast(node->input(kOneNum)); @@ -3032,7 +3123,7 @@ void OnnxExporter::ExportMergeMaxPoolWithArgmax(const FuncGraphPtr &func_graph, // LayerNorm(N, C1, H, W) --> reshape(1, C2, 1, W) + MeanVarianceNormalization + reshape(N, C1, H, W) void OnnxExporter::ExportMergeLayerNorm(const FuncGraphPtr &, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto LayerNormNode = dyn_cast(node->input(kOneNum)); auto layernorm_input_x = GetNodeInputName(LayerNormNode->input(kOneNum), node_map_ptr, graph_proto); @@ -3047,17 +3138,16 @@ void OnnxExporter::ExportMergeLayerNorm(const FuncGraphPtr &, const CNodePtr &no auto onnx_type = GetOutputType(LayerNormNode->input(kOneNum)); auto input_shape = dyn_cast(LayerNormNode->input(kOneNum)->Shape())->shape(); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); auto epsilon = GetOpAttribute(LayerNormNode, "epsilon"); std::vector reduce_axes = {static_cast(input_shape.size()) - 1}; - AddMeanVarianceNormalizationOp(layernorm_input_x, layernorm_input_gamma, layernorm_input_beta, - std::to_string(node_idx), reduce_axes, epsilon, input_shape, onnx_type, graph_proto); + AddMeanVarianceNormalizationOp(layernorm_input_x, layernorm_input_gamma, layernorm_input_beta, node_name, reduce_axes, + epsilon, input_shape, onnx_type, graph_proto); } void OnnxExporter::ExportMergeConv2DTranspose(const FuncGraphPtr &func_graph, const CNodePtr &node, - std::map *node_map_ptr, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto conv_node = dyn_cast(node->input(kOneNum)); PrimConv2DTransposeExportHelper(conv_node, node, node_map_ptr, graph_proto); @@ -3081,13 +3171,9 @@ void OnnxExporter::ExportMergeConv2DTranspose(const FuncGraphPtr &func_graph, co return self.x, self.x */ -void OnnxExporter::ExportOutput(const FuncGraphPtr &, const CNodePtr &node, std::map *node_map_ptr, - onnx::GraphProto *const graph_proto) { - if (node->inputs().size() != kTwoNum) { - MS_LOG(EXCEPTION) << "Number of inputs of return node is not equal to 2."; - } - - AnfNodePtr arg = GetRealInput(node->input(kOneNum)); +void OnnxExporter::ExportOutput(const FuncGraphPtr &, const AnfNodePtr &return_arg, + std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { + AnfNodePtr arg = GetRealInput(return_arg); if (IsPrimitiveCNode(arg, prim::kPrimMakeTuple)) { auto arg_cnode = dyn_cast(arg); for (size_t i = 1; i < arg_cnode->inputs().size(); ++i) { @@ -3102,9 +3188,7 @@ void OnnxExporter::ExportOutput(const FuncGraphPtr &, const CNodePtr &node, std: auto tuple = arg->cast()->value()->cast(); for (size_t i = 0; i < tuple->value().size(); ++i) { const auto &element = tuple->value().at(i); - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; - std::string output_name = std::to_string(node_idx); + std::string output_name = GenerateUniqueName(); onnx::TensorProto *initializer = graph_proto->add_initializer(); initializer->set_name(output_name); @@ -3135,45 +3219,46 @@ void OnnxExporter::ExportOutput(const FuncGraphPtr &, const CNodePtr &node, std: } } -std::string OnnxExporter::GetNodeInputName(const AnfNodePtr &orig_node, std::map *node_map_ptr, +std::string OnnxExporter::GetNodeInputName(const AnfNodePtr &orig_node, std::map *node_map_ptr, onnx::GraphProto *const graph_proto) { auto node = GetRealInput(orig_node); - if (node->isa()) { - auto iter = node_map_ptr->find(node); - if (iter == node_map_ptr->end()) { - MS_LOG(EXCEPTION) << "Can not find node '" << node->DebugString() << "' in node_map"; - } - return std::to_string(iter->second); + + auto renamed_iter = renamed_node_map_.find(node); + if (renamed_iter != renamed_node_map_.end()) { + return renamed_iter->second; } - if (node->isa() && !node->cast()->has_default()) { - return node->ToString(); + auto iter = node_map_ptr->find(node); + if (iter != node_map_ptr->end()) { + return iter->second; + } + + if (node->isa() || (node->isa() && !node->cast()->has_default())) { + MS_LOG(EXCEPTION) << "Can not find node '" << node->DebugString() << "' in node_map"; } // for ValueNode or Parameter with default input, create an initializer - if (node->isa() || node->isa()) { - auto iter = node_map_ptr->find(node); - if (iter != node_map_ptr->end()) { - return std::to_string(iter->second); - } - // the id number starts at 1, so the id of created node should be size of map plus one - auto node_idx = AllocateNodeIndex(); - (*node_map_ptr)[node] = node_idx; - std::string node_name = std::to_string(node_idx); - - ValuePtr value; - if (node->isa()) { - value = node->cast()->value(); - } else if (node->isa()) { - value = node->cast()->default_param(); - } else { - MS_LOG(EXCEPTION) << "Impossible branch"; - } + if (node->isa()) { + auto node_name = RegisterNodeWithUniqueName(node, node_map_ptr); + auto value = node->cast()->value(); onnx::TensorProto *initializer_proto = graph_proto->add_initializer(); initializer_proto->set_name(node_name); SetTensorData(value, initializer_proto); + (*node_map_ptr)[node] = node_name; + return node_name; + } + + if (node->isa()) { + auto param = dyn_cast(node); + auto node_name = GenerateUniqueParameterName(param, node_map_ptr); + + onnx::TensorProto *initializer_proto = graph_proto->add_initializer(); + initializer_proto->set_name(node_name); + SetTensorData(param->default_param(), initializer_proto); + + (*node_map_ptr)[node] = node_name; return node_name; }