From eaaecac92cf7f86f53e57b16a9124c7e06dfdcde Mon Sep 17 00:00:00 2001 From: zyf1234 Date: Tue, 5 Sep 2023 22:08:38 +0800 Subject: [PATCH] ADD file via upload --- .../ccsrc/transform-update/mindir_exporter.cc | 1311 +++++++++++++++++ 1 file changed, 1311 insertions(+) create mode 100644 mindspore/ccsrc/transform-update/mindir_exporter.cc diff --git a/mindspore/ccsrc/transform-update/mindir_exporter.cc b/mindspore/ccsrc/transform-update/mindir_exporter.cc new file mode 100644 index 00000000000..a2e6b397df6 --- /dev/null +++ b/mindspore/ccsrc/transform-update/mindir_exporter.cc @@ -0,0 +1,1311 @@ +/** + * Copyright 2020-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 +#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 "mindspore/core/ops/core_ops.h" +#include "proto/mind_ir.pb.h" +#include "utils/check_convert_utils.h" +#include "include/common/debug/dump_proto.h" +#include "utils/ms_utils.h" +#include "include/common/utils/utils.h" +#ifndef MINDIR_EXPORT_TENSOR_LAYOUT_CLIP +#include "frontend/parallel/tensor_layout/tensor_layout.h" +#endif +#include "abstract/abstract_function.h" +#include "mindspore/core/utils/file_utils.h" + +namespace mindspore { +using FloatPtr = std::shared_ptr; +using IntPtr = std::shared_ptr; +using UIntPtr = std::shared_ptr; +using ModelProtoPtr = std::shared_ptr; + +// anf type to mindir type map将 ANF 类型映射到 MindIR 类型的映射表 +static mindspore::HashMap g_data_type_map = { + {kNumberTypeBool, mind_ir::TensorProto_DataType_BOOL}, + {kNumberTypeInt8, mind_ir::TensorProto_DataType_INT8}, + {kNumberTypeInt16, mind_ir::TensorProto_DataType_INT16}, + {kNumberTypeInt32, mind_ir::TensorProto_DataType_INT32}, + {kNumberTypeInt64, mind_ir::TensorProto_DataType_INT64}, + {kNumberTypeUInt8, mind_ir::TensorProto_DataType_UINT8}, + {kNumberTypeUInt16, mind_ir::TensorProto_DataType_UINT16}, + {kNumberTypeUInt32, mind_ir::TensorProto_DataType_UINT32}, + {kNumberTypeUInt64, mind_ir::TensorProto_DataType_UINT64}, + {kNumberTypeFloat16, mind_ir::TensorProto_DataType_FLOAT16}, + {kNumberTypeFloat32, mind_ir::TensorProto_DataType_FLOAT}, + {kNumberTypeFloat64, mind_ir::TensorProto_DataType_DOUBLE}, + {kObjectTypeString, mind_ir::TensorProto_DataType_STRING}, + {kNumberTypeComplex64, mind_ir::TensorProto_DataType_COMPLEX64}, + {kNumberTypeComplex128, mind_ir::TensorProto_DataType_COMPLEX128}}; + +static mindspore::HashMap g_data_bits_int_map = { + {8, mind_ir::TensorProto_DataType_INT8}, + {16, mind_ir::TensorProto_DataType_INT16}, + {32, mind_ir::TensorProto_DataType_INT32}, + {64, mind_ir::TensorProto_DataType_INT64}, +}; + +static mindspore::HashMap g_data_bits_uint_map = { + {8, mind_ir::TensorProto_DataType_UINT8}, + {16, mind_ir::TensorProto_DataType_UINT16}, + {32, mind_ir::TensorProto_DataType_UINT32}, + {64, mind_ir::TensorProto_DataType_UINT64}, +}; + +static mindspore::HashMap g_data_bits_float_map = { + {16, mind_ir::TensorProto_DataType_FLOAT16}, + {32, mind_ir::TensorProto_DataType_FLOAT}, + {64, mind_ir::TensorProto_DataType_FLOAT64}, +}; + +static std::set g_export_attr_blacklist = {kAttrDump}; + +// Can build different builder according to format根据格式构建不同的生成器。 +class IrExportBuilder; +using IrExportBuilderPtr = std::shared_ptr; +//使用IrExportBuilderPtr表示std::shared_ptr + +class IrExporter { + public: + explicit IrExporter(IrExportBuilderPtr builder) : builder_(std::move(builder)) {} + virtual ~IrExporter() = default; + std::string GetDumpString(const FuncGraphPtr &func_graph); + ModelProtoPtr GetDumpProto(const FuncGraphPtr &func_graph, const FuncGraphPtr ¶m_layout_fg = nullptr); + + private: + IrExportBuilderPtr builder_; +}; +//声明IrExporter类 +using IrExporterPtr = std::shared_ptr; + +class IrExportBuilder { + public: + IrExportBuilder() : model_(std::make_shared()) {} + ~IrExportBuilder() = default; + std::string GetProtoString() const; + void BuildModelInfo(); + bool BuildModel(const FuncGraphPtr &func_graph); + ModelProtoPtr Model() { return model_; } + +#ifndef MINDIR_EXPORT_TENSOR_LAYOUT_CLIP + void BuildLayout(const FuncGraphPtr &func_graph); +#endif +//如果没有定义MINDIR_EXPORT_TENSOR_LAYOUT_CLIP,则会定义一个名为BuildLayout的函数 + + bool BuildFuncGraph(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto); + bool BuildFuncGraphAttrs(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto); + bool BuildParameters(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto); + bool BuildNodes(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto); + bool BuildOutput(const CNodePtr &node, mind_ir::GraphProto *const graph_proto); + bool BuildCNode(const CNodePtr &node, mind_ir::GraphProto *const graph_proto); + std::string BuildInputNode(const AnfNodePtr &node, mind_ir::GraphProto *const graph_proto); + + bool SetValueInfoProto(const AnfNodePtr &node, mind_ir::ValueInfoProto *const value_proto); + bool SetParamToTensorProto(const ParameterPtr ¶m, mind_ir::TensorProto *const tensor_proto); + bool SetTensorProto(const AbstractBasePtr &abstract, mind_ir::TensorProto *const tensor_proto); + bool SetCSRTensorToProto(const AbstractBasePtr &abstract, mind_ir::AttributeProto *const attr_proto); + bool SetCOOTensorToProto(const AbstractBasePtr &abstract, mind_ir::AttributeProto *const attr_proto); + bool SetAttributeProto(const AnfNodePtr &node, mind_ir::NodeProto *const node_proto); + bool SetAbstractToNodeProto(const CNodePtr &node, mind_ir::NodeProto *const node_proto); + bool SetAbstractToNodeProto(const abstract::AbstractBasePtr &abstract, mind_ir::AttributeProto *const attr_proto); + bool SetValueToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto); + bool SetTypeToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto); + bool SetScalarToAttributeProto_ir(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) const; + bool SetScalarToAttributeProtoForInt_ir(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) const; + bool SetScalarToAttributeProto_irs(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) const; + bool SetScalarToAttributeProtoForInt_irs(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) const; + bool SetTypeToAttributeProto_irs(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto); + bool SetTensorToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto); + bool SetSequenceToAttributeProto(const ValueSequencePtr &value, mind_ir::AttributeProto *const attr_proto); + bool SetSeqElemToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto); + + mind_ir::TensorProto_DataType GetMindirDataType(TypeId type_id) const; + mind_ir::TensorProto_DataType GetMindirDataBitsIntType(int bits) const; + mind_ir::TensorProto_DataType GetMindirDataBitsFloatType(int bits) const; + mind_ir::TensorProto_DataType GetMindirDataBitsUIntType(int bits) const; + std::string GetNodeName(const AnfNodePtr &node) const; + std::string GetUniqueNodeName(const AnfNodePtr &node); + std::string GetOpTypeName(const AnfNodePtr &node); + size_t GetUniqueID() { return ++unique_id_; } + + private: + bool SetAbstractFuncToAttributeProto(const abstract::AbstractBasePtr &abstract, + mind_ir::AttributeProto *const attr_proto); + std::string GetPrimitiveUniqueName(const PrimitivePtr &primitive_ptr); + bool BuildPrimitives(); + + ModelProtoPtr model_; + mind_ir::NodeProto *last_node_{nullptr}; + std::list todo_; + std::map node_name_map_; + std::map primitive_name_map_; + std::set nodeName_; + size_t unique_id_{0}; + bool top_graph{true}; +}; +//声明IrExportBuilder类 + +bool IrExportBuilder::SetAbstractFuncToAttributeProto(const abstract::AbstractBasePtr &abstract, + mind_ir::AttributeProto *const attr_proto) { + MS_EXCEPTION_IF_NULL(abstract); + MS_EXCEPTION_IF_NULL(attr_proto); + //如果abstract、attr_proto为空,则抛出异常 + if (abstract->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_FUNCGRAPHCLOSURE); + auto func_name = abstract->cast()->func_graph()->ToString(); + attr_proto->set_s(func_name); + } else if (abstract->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_PRIMITIVECLOSURE); + auto prim = abstract->cast()->prim(); + attr_proto->set_s(GetPrimitiveUniqueName(prim)); + } else if (abstract->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_PARTIALCLOSURE); + auto node_ptr = abstract->cast()->node(); + MS_EXCEPTION_IF_NULL(node_ptr); + attr_proto->set_s(GetUniqueNodeName(node_ptr)); + } else if (abstract->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UNIONFUNCCLOSURE); + auto visit_func = [this, &attr_proto](const abstract::AbstractFuncAtomPtr &poss) { + auto element_attr_proto = attr_proto->add_values(); + if (!this->SetAbstractFuncToAttributeProto(poss, element_attr_proto)) { + MS_LOG(EXCEPTION) << "Set union function abstract to proto error." << poss->ToString(); + } + }; + abstract->cast()->Visit(visit_func); + } else { + MS_LOG(ERROR) << "The parameter abstract is not an abstractFunction: " << abstract->ToString(); + return false; + } + return true; +} +//实例化IrExportBuilder类中的函数SetAbstractFuncToAttributeProto +//根据abstract变量下的各项事例是否存在,设置attr_proto的类型为对应名称,然后从abstract变量中获取对应类型的指针,进而获取其对应的函数图名称 +//并将该名称设置为 attr_proto的字符串属性,并返回为真。若没有对应事例则抛出异常并报错,返回为假。 + +std::string IrExportBuilder::GetPrimitiveUniqueName(const PrimitivePtr &primitive_ptr) { + auto it = primitive_name_map_.find(primitive_ptr); + if (it != primitive_name_map_.end()) { + return it->second; + } + // Remove this check if we find a way to handle save/load training model with flattened parameters. + if (IsPrimitiveEquals(primitive_ptr, prim::kPrimFlattenConcat)) { + MS_LOG(EXCEPTION) << "Export model with operator '" << primitive_ptr->name() << "' is not supported yet.\n" + << "Please remove 'net.flatten_weights()' in your script and try again."; + } + auto answer = primitive_ptr->name() + ":" + std::to_string(GetUniqueID()); + primitive_name_map_[primitive_ptr] = answer; + return answer; +} +//实例化IrExportBuilder类中的函数GetPrimitiveUniqueName +//如果找到一种方法来处理具有扁平参数的保存/加载训练模型,请删除此检查 + +bool IrExportBuilder::BuildPrimitives() { + // 遍历 primitive_name_map_ 中的每个原语 + for (auto it = primitive_name_map_.begin(); it != primitive_name_map_.end(); ++it) { + auto prim_proto = model_->add_primitives(); + auto prim = it->first; + prim_proto->set_name(it->second); + prim_proto->set_op_type(prim->name()); + // 获取实际的原语(可能存在原语的包装) + auto real_prim = GetValueWithoutDoSignature(prim)->cast(); + if (real_prim != nullptr) { + prim = real_prim; + } + + // Set primitive attributes遍历设置原语的属性 + for (const auto &attr : prim->attrs()) { + // 检查当前属性是否在黑名单中,如果是则跳过 + MS_LOG(DEBUG) << "attr: " << attr.first << " " << attr.second->DumpText() << " " << attr.second->type_name(); + auto iter = g_export_attr_blacklist.find(attr.first); + if (iter != g_export_attr_blacklist.end()) { + continue; + } + // 向原语的 proto 中添加一个属性 + mind_ir::AttributeProto *attr_proto = prim_proto->add_attribute(); + attr_proto->set_name(attr.first); + auto attr_value = attr.second; + // 转换并检查属性值 + CheckAndConvertUtils::ConvertAttrValueInExport(prim->name(), attr.first, &attr_value); + if (!SetValueToAttributeProto(attr_value, attr_proto)) { + MS_LOG(ERROR) << "Set value to AttributeProto failed."; + return false; + } + } // Loop of attrs + } // Loop of primitives + return true; +} + +std::string IrExporter::GetDumpString(const FuncGraphPtr &func_graph) { + auto dump_proto = GetDumpProto(func_graph); + if (dump_proto == nullptr) { + MS_LOG(EXCEPTION) << "Get dump proto for graph " << func_graph->ToString() << " failed."; + } + return builder_->GetProtoString(); +} +//获取DumpString + +ModelProtoPtr IrExporter::GetDumpProto(const FuncGraphPtr &func_graph, const FuncGraphPtr ¶m_layout_fg) { + if ((builder_ == nullptr) || (func_graph == nullptr)) { + MS_LOG(EXCEPTION) << "Input params is null."; + } + + // Export model info + builder_->BuildModelInfo(); + + // Export model and return string + if (!builder_->BuildModel(func_graph)) { + return nullptr; + } + +#ifndef MINDIR_EXPORT_TENSOR_LAYOUT_CLIP + // Export layout information + if (param_layout_fg) { + builder_->BuildLayout(param_layout_fg); + } +#endif + return builder_->Model(); +} + +std::string IrExportBuilder::GetProtoString() const { + MS_LOG(DEBUG) << "BuildModel complete!"; + return model_->SerializeAsString(); +} + +void IrExportBuilder::BuildModelInfo() { + // 构建模型信息 + constexpr auto ir_version = "0.1.1"; + constexpr auto mindspore_name = "MindSpore"; + model_->set_ir_version(ir_version);// 设置IR版本 + model_->set_producer_name(mindspore_name);// 设置生产者名称 + model_->set_model_version(VERSION);// 设置模型版本 + model_->set_little_endian(common::IsLittleByteOrder());// 设置字节序 + model_->set_mind_ir_version(mind_ir::Version_MAX);// 设置Mind IR版本 +} + +#ifndef MINDIR_EXPORT_TENSOR_LAYOUT_CLIP +void IrExportBuilder::BuildLayout(const FuncGraphPtr &func_graph) { + // 构建张量布局信息 + MS_EXCEPTION_IF_NULL(func_graph); + std::vector graph_params = func_graph->parameters();// 获取图的参数节点 + mind_ir::ParallelProto *parallel_proto = model_->mutable_parallel();// 获取模型的并行信息 + // 遍历图的参数节点 + for (auto para : graph_params) { + std::string name = std::static_pointer_cast(para)->name();// 获取参数节点的名称 + auto tensor_layout = para->user_data();// 获取参数节点的张量布局信息 + if (tensor_layout == nullptr) { + MS_LOG(INFO) << "GetParameterLayout nullptr name = " << name; + } else { + mind_ir::LayoutProto *layoutProto = parallel_proto->add_layout();// 添加张量布局信息到模型的并行信息中 + + // Get all the information for layput + // 获取张量布局的各种信息 + auto device_arrangement = tensor_layout->device_arrangement().array(); + auto tensor_map = tensor_layout->tensor_map().array(); + auto slice_shape = tensor_layout->slice_shape().array(); + int64_t field_size = tensor_layout->get_field_size(); + bool uniform_split = tensor_layout->uniform_split(); + std::string opt_shard_group = tensor_layout->opt_shard_group(); + + // Save all information to Layout Proto + // 将信息保存到布局信息中 + layoutProto->set_name(name); + for (auto device_arrangement_element : device_arrangement) { + layoutProto->add_device_arrangement_int(device_arrangement_element); + } + for (auto tensor_map_element : tensor_map) { + layoutProto->add_tensor_map_int(tensor_map_element); + } + for (auto slice_shape_element : slice_shape) { + layoutProto->add_slice_shape_int(slice_shape_element); + } + layoutProto->set_field_size(field_size); + layoutProto->set_uniform_split(uniform_split); + layoutProto->set_opt_shard_group(opt_shard_group); + } + } +} +#endif + +bool IrExportBuilder::BuildModel(const FuncGraphPtr &func_graph) { + // 构建模型的函数 + MS_EXCEPTION_IF_NULL(func_graph);// 检查输入函数图是否为空 + // 清空待办列表、节点名称集合和原语名称映射 + mind_ir::GraphProto *graph_proto = model_->mutable_graph(); + graph_proto->set_name(func_graph->ToString()); + graph_proto->set_bprop_hash(func_graph->bprop_hash()); + // 清空待办列表、节点名称集合和原语名称映射 + todo_.clear(); + nodeName_.clear(); + primitive_name_map_.clear(); + // Build the main funcGraph + // 构建主函数图 + // 将主函数图名称添加到节点名称集合 + (void)nodeName_.insert(func_graph->ToString()); + top_graph = true; + if (!BuildFuncGraph(func_graph, graph_proto)) { + MS_LOG(ERROR) << "Build func_graph " << func_graph->ToString() << " failed."; + return false; + } + + // Build child funcGraphs + // 构建子函数图 + std::set graphVisited; + (void)graphVisited.insert(func_graph); + top_graph = false; + while (!todo_.empty()) { + // 从待办列表中取出一个函数图 + FuncGraphPtr fg = todo_.back(); + todo_.pop_back(); + // 如果函数图已经被访问过,则继续处理下一个函数图 + if (graphVisited.count(fg) > 0) { + continue; + } + // 检查节点名称是否重复,如果重复则报错 + if (nodeName_.count(fg->ToString()) > 0) { + MS_LOG(ERROR) << "There is a duplicate name: " << fg->ToString(); + return false; + } + // 将函数图名称添加到节点名称集合和已访问的函数图集合 + (void)nodeName_.insert(fg->ToString()); + (void)graphVisited.insert(fg); + // 创建一个新的函数图对象,并构建该函数图 + auto graph = model_->add_functions(); + if (!BuildFuncGraph(fg, graph)) { + MS_LOG(ERROR) << "Build func_graph " << fg->ToString() << " failed."; + return false; + } + } + // 构建原语信息 + if (!BuildPrimitives()) { + return false; + } + // Release resource + // 释放资源,清空节点名称集合、节点名称映射和原语名称映射 + nodeName_.clear(); + node_name_map_.clear(); + primitive_name_map_.clear(); + return true; +} + +bool IrExportBuilder::BuildFuncGraph(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto) { + // Export funcGraph name. + graph_proto->set_name(func_graph->ToString()); + // Export parameters + // 1. parameters should be mapped to ValueInfoProto + // 2. parameters with default value should be mapped to Initializer + //导出参数 + //1.参数应映射到ValueInfoProto + //2.具有默认值的参数应映射到Initializer + if (!BuildParameters(func_graph, graph_proto)) { + MS_LOG(ERROR) << "Build parameters failed."; + return false; + } + + // Export graph attributes + //导出图形属性 + if (!BuildFuncGraphAttrs(func_graph, graph_proto)) { + MS_LOG(ERROR) << "Build attributes for graph failed."; + return false; + } + + // Export operator nodes(include output) + //导出操作员节点(包括输出) + return BuildNodes(func_graph, graph_proto); +} + +bool IrExportBuilder::BuildFuncGraphAttrs(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto) { + MS_EXCEPTION_IF_NULL(func_graph); + MS_EXCEPTION_IF_NULL(graph_proto); + // 遍历函数图的所有属性 + for (const auto &attr : func_graph->attrs()) { + // 输出调试信息,打印属性名、属性值的文本表示和属性值的类型名 + MS_LOG(DEBUG) << "attr: " << attr.first << " " << attr.second->DumpText() << " " << attr.second->type_name(); + // 在导出属性黑名单中查找当前属性名 + auto iter = g_export_attr_blacklist.find(attr.first); + if (iter != g_export_attr_blacklist.end()) { + continue; + } + // 创建一个AttributeProto对象,并设置属性名称 + mind_ir::AttributeProto *attr_proto = graph_proto->add_attribute(); + attr_proto->set_name(attr.first); + // 将属性值转换并设置到AttributeProto中 + if (!SetValueToAttributeProto(attr.second, attr_proto)) { + MS_LOG(ERROR) << "Set value to AttributeProto for GraphProto failed."; + return false; + } + } + return true; +} + +bool IrExportBuilder::BuildParameters(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto) { + // 构建函数图的参数信息并添加到GraphProto中 + MS_EXCEPTION_IF_NULL(func_graph); + MS_EXCEPTION_IF_NULL(graph_proto); + // 遍历函数图的所有参数节点 + for (auto &item : func_graph->parameters()) { + MS_EXCEPTION_IF_NULL(item); + auto param = item->cast(); + // 如果无法将节点转换为参数节点,输出错误信息并返回失败 + if (param == nullptr) { + MS_LOG(ERROR) << "Parameter: '" << item->ToString() << "' could not cast to parameter."; + return false; + } + // 获取唯一的参数名称 + std::string param_name = GetUniqueNodeName(param); + // 如果是顶层函数图且参数具有默认值 + if (top_graph && param->has_default()) { + MS_LOG(DEBUG) << "Parameter: '" << item->DebugString(); + mind_ir::TensorProto *parameter_proto = graph_proto->add_parameter(); + // 设置参数节点的名称,并将参数转换为TensorProto + parameter_proto->set_name(param_name); + if (!SetParamToTensorProto(param, parameter_proto)) { + MS_LOG(ERROR) << "Set parameter " << param->DebugString() << " to TensorProto failed."; + return false; + } + } else { + mind_ir::ValueInfoProto *input_proto = graph_proto->add_input(); + // 设置参数节点的名称,并将参数转换为ValueInfoProto + input_proto->set_name(param_name); + // 检查参数名称是否重复,如果是则输出错误信息并返回失败 + if (!SetValueInfoProto(param, input_proto)) { + MS_LOG(ERROR) << "Set parameter " << param->DebugString() << " to TensorProto failed."; + return false; + } + } + if (nodeName_.count(param_name) > 0) { + MS_LOG(ERROR) << "parameter name is duplicate:" << param_name; + return false; + } + (void)nodeName_.insert(param_name); + } + return true; +} + +mind_ir::TensorProto_DataType IrExportBuilder::GetMindirDataType(TypeId type_id) const { + auto iter = g_data_type_map.find(type_id); + if (iter == g_data_type_map.end()) { + MS_LOG(ERROR) << "Convert type error, unsupported type! " << type_id; + return mind_ir::TensorProto_DataType_UNDEFINED; + } + return iter->second; +} +//获取Mindir数据类型 + +mind_ir::TensorProto_DataType IrExportBuilder::GetMindirDataBitsIntType(int bits) const { + auto iter = g_data_bits_int_map.find(bits); + if (iter == g_data_bits_int_map.end()) { + MS_LOG(ERROR) << "Convert bits int error, unsupported bits! " << bits; + return mind_ir::TensorProto_DataType_UNDEFINED; + } + return iter->second; +} +//获取Mindir数据是否为int类型 + +mind_ir::TensorProto_DataType IrExportBuilder::GetMindirDataBitsUIntType(int bits) const { + auto iter = g_data_bits_uint_map.find(bits); + if (iter == g_data_bits_uint_map.end()) { + MS_LOG(ERROR) << "Convert bits uint error, unsupported bits! " << bits; + return mind_ir::TensorProto_DataType_UNDEFINED; + } + return iter->second; +} +//获取Mindir数据是否为uint类型 +mind_ir::TensorProto_DataType IrExportBuilder::GetMindirDataBitsFloatType(int bits) const { + auto iter = g_data_bits_float_map.find(bits); + if (iter == g_data_bits_float_map.end()) { + MS_LOG(ERROR) << "Convert bits float error, unsupported bits! " << bits; + return mind_ir::TensorProto_DataType_UNDEFINED; + } + return iter->second; +} +//获取Mindir数据是否为float类型 +bool IrExportBuilder::SetValueInfoProto(const AnfNodePtr &node, mind_ir::ValueInfoProto *const value_proto) { + if (node == nullptr || value_proto == nullptr) { + MS_LOG(EXCEPTION) << "AnfNode or ValueInfo is null!"; + } + MS_LOG(DEBUG) << "SetValueInfoProto: " << node->DebugString(); + const TypePtr &type = node->Type(); + const BaseShapePtr &shape = node->Shape(); + // For the bprop fg which has not been renormalized. + if (type == nullptr || shape == nullptr) { + return true; + } + if (type->isa() && shape->isa()) { + mind_ir::TensorProto *tensor_proto = value_proto->add_tensor(); + if (!SetTensorProto(node->abstract(), tensor_proto)) { + return false; + } + } else if (type->isa()) { + mind_ir::AttributeProto *attribute = value_proto->mutable_attr_info(); + if (!SetAbstractToNodeProto(node->abstract(), attribute)) { + MS_LOG(ERROR) << "Set shape to Proto for " << node->DebugString() << " failed."; + return false; + } + attribute->set_name("shape"); + } else { + value_proto->set_denotation(type->type_name()); + } + MS_LOG(DEBUG) << "Value type: " << type->type_name(); + return true; +} +//设置proto参数 + +bool IrExportBuilder::SetTensorToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) { + if (value == nullptr || attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "ValuePtr or AttributeProto is null!"; + } + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + tensor_proto->set_name("value0"); + auto data = value->cast(); + MS_EXCEPTION_IF_NULL(data); + tensor_proto->set_raw_data(data->data_c(), static_cast(data->data().nbytes())); + auto dtype = data->data_type(); + auto shape = data->shape_c(); + auto data_type = GetMindirDataType(dtype); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + for (const auto &dim : shape) { + tensor_proto->add_dims(dim); + } + return true; +} +//设置proto参数 +bool IrExportBuilder::SetCSRTensorToProto(const AbstractBasePtr &abstract, mind_ir::AttributeProto *const attr_proto) { + abstract::AbstractCSRTensorPtr csr_tensor_abs = abstract->cast(); + MS_EXCEPTION_IF_NULL(csr_tensor_abs); + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_CSR_TENSOR); + mind_ir::AttributeProto *indptr = attr_proto->add_values(); + bool res = SetAbstractToNodeProto(csr_tensor_abs->indptr(), indptr); + mind_ir::AttributeProto *indices = attr_proto->add_values(); + res = res && SetAbstractToNodeProto(csr_tensor_abs->indices(), indices); + mind_ir::AttributeProto *values = attr_proto->add_values(); + res = res && SetAbstractToNodeProto(csr_tensor_abs->values(), values); + mind_ir::AttributeProto *shape = attr_proto->add_values(); + res = res && SetAbstractToNodeProto(csr_tensor_abs->shape(), shape); + return res; +} +//设置proto参数 +bool IrExportBuilder::SetCOOTensorToProto(const AbstractBasePtr &abstract, mind_ir::AttributeProto *const attr_proto) { + abstract::AbstractCOOTensorPtr coo_tensor_abs = abstract->cast(); + MS_EXCEPTION_IF_NULL(coo_tensor_abs); + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_COO_TENSOR); + mind_ir::AttributeProto *indices = attr_proto->add_values(); + bool res = SetAbstractToNodeProto(coo_tensor_abs->indices(), indices); + mind_ir::AttributeProto *values = attr_proto->add_values(); + res = res && SetAbstractToNodeProto(coo_tensor_abs->values(), values); + mind_ir::AttributeProto *shape = attr_proto->add_values(); + res = res && SetAbstractToNodeProto(coo_tensor_abs->shape(), shape); + return res; +} +//设置proto参数 +bool IrExportBuilder::SetTensorProto(const AbstractBasePtr &abstract, mind_ir::TensorProto *const tensor_proto) { + auto type = abstract->BuildType(); + auto shape = abstract->BuildShape(); + if (!type->isa() || !shape->isa()) { + MS_LOG(ERROR) << "Type or shape is not supported! " << type->ToString(); + return false; + } + auto tensor = type->cast(); + auto tensor_shape = shape->cast(); + const auto &dims = tensor_shape->shape(); + auto data_type = GetMindirDataType(tensor->element()->type_id()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + for (const auto &dim : dims) { + tensor_proto->add_dims(dim); + } + if (tensor_shape->IsDynamic()) { + auto min_shape = tensor_shape->min_shape(); + auto max_shape = tensor_shape->max_shape(); + for (auto item : min_shape) { + tensor_proto->add_min_dims(item); + } + for (auto item : max_shape) { + tensor_proto->add_max_dims(item); + } + } + if (!abstract->name().empty()) { + tensor_proto->set_name(abstract->name()); + } + // Deal Ref + if (!type->isa()) { + return true; + } + + auto abs_ref = abstract->cast(); + if (abs_ref == nullptr) { + MS_LOG(ERROR) << "The abstract " << abstract->ToString() << " should be AbstractRefTensor."; + return false; + } + auto ref_key_value = abs_ref->ref_key_value()->cast(); + if (ref_key_value == nullptr) { + MS_LOG(INFO) << "The ref_key_value of abstract ref " << abstract->ToString() << " is nullptr"; + return true; + } + tensor_proto->set_ref_key(ref_key_value->value()); + return true; +} +//设置proto参数 +bool IrExportBuilder::SetParamToTensorProto(const ParameterPtr ¶m, mind_ir::TensorProto *const tensor_proto) { + if (param == nullptr || tensor_proto == nullptr) { + MS_LOG(EXCEPTION) << "Parameter or TensorProto is null!"; + } + MS_LOG(DEBUG) << "SetParamToTensorProto: " << param->DebugString(); + return SetTensorProto(param->abstract(), tensor_proto); +} +//设置proto参数 +bool IrExportBuilder::BuildNodes(const FuncGraphPtr &func_graph, mind_ir::GraphProto *const graph_proto) { + // 构建函数图中的节点信息并添加到GraphProto中 + std::vector nodes = TopoSort(func_graph->get_return(), SuccIncoming, AlwaysInclude);// 使用拓扑排序获取函数图中的节点顺序 + for (const AnfNodePtr &node : nodes) {// 遍历所有节点 + MS_EXCEPTION_IF_NULL(node); + // 如果节点不是CNode类型,则输出调试信息并继续处理下一个节点 + if (!node->isa()) { + MS_LOG(DEBUG) << "Node: '" << node->ToString() << "' is not cnode"; + continue; + } + auto cnode = node->cast(); + // 如果节点是函数图的返回节点 + if (cnode == func_graph->get_return()) { + // 构建返回节点的输出信息并添加到GraphProto + if (!BuildOutput(cnode, graph_proto)) { + MS_LOG(ERROR) << "Build output for graph " << func_graph->ToString() << " failed."; + return false; + } + } else { + // 构建普通CNode节点的信息并添加到GraphProto + if (!BuildCNode(cnode, graph_proto)) { + MS_LOG(ERROR) << "Build proto for cnode " << cnode->DebugString() << " failed."; + return false; + } + } + } + return true; +} + +bool IrExportBuilder::BuildOutput(const CNodePtr &node, mind_ir::GraphProto *const graph_proto) { + MS_EXCEPTION_IF_NULL(node); + const int OutputSize = 2; + if (node->size() != OutputSize) { + MS_LOG(ERROR) << "Number of inputs of return node is not equal to 2."; + return false; + } + AnfNodePtr arg = node->input(1); + std::string node_name = BuildInputNode(arg, graph_proto); + if (node_name.empty()) { + MS_LOG(ERROR) << "Build input node failed for arg " << arg->DebugString(); + return false; + } + mind_ir::ValueInfoProto *output_proto = graph_proto->add_output(); + output_proto->set_name(node_name); + return SetValueInfoProto(arg, output_proto); +} +//建立输出 +std::string IrExportBuilder::GetOpTypeName(const AnfNodePtr &node) { + // May be ValueNode/CNode/Parameter + std::string type_name = ""; + if (IsValueNode(node)) { + PrimitivePtr prim = GetValueNode(node); + MS_EXCEPTION_IF_NULL(prim); + type_name = "REF::" + GetPrimitiveUniqueName(prim); + } else if (IsValueNode(node)) { + FuncGraphPtr fg = GetValueNode(node); + MS_EXCEPTION_IF_NULL(fg); + todo_.push_back(fg); + type_name = "REF::" + fg->ToString(); + } else if (node->isa() || node->isa()) { + auto nodeName = GetUniqueNodeName(node); + type_name = "REF::" + nodeName; + if (nodeName_.count(nodeName) == 0) { + MS_LOG(ERROR) << "There is not the name: " << nodeName; + return ""; + } + } else { + MS_LOG(ERROR) << "Need to support op type: " << node->type_name(); + return ""; + } + MS_LOG(DEBUG) << "ExportType: " << type_name; + return type_name; +} +//获取OpType的类型名,可能为ValueNode/CNode/Parameter +bool IrExportBuilder::SetAbstractToNodeProto(const AbstractBasePtr &abs, mind_ir::AttributeProto *const attr_proto) { + auto type = abs->BuildType(); + auto shape = abs->BuildShape(); + if (type->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TUPLE); + auto tuple_abs = abs->cast(); + for (size_t i = 0; i < tuple_abs->size(); i++) { + mind_ir::AttributeProto *attr_values = attr_proto->add_values(); + if (!SetAbstractToNodeProto((*tuple_abs)[i], attr_values)) { + return false; + } + } + } else if (type->isa() && shape->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + return SetTensorProto(abs, tensor_proto); + } else if (type->isa()) { + if (type->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_BOOL); + } else { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + auto data_type = GetMindirDataType(type->type_id()); + tensor_proto->set_data_type(data_type); + tensor_proto->add_dims(1); + } + } else if (type->isa()) { + if (!SetAbstractFuncToAttributeProto(abs, attr_proto)) { + return false; + } + } else if (type->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_STRING); + } else if (type->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UMONAD); + } else if (type->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_IOMONAD); + } else if (type->isa()) { + auto csr_tensor_abs = abs->cast(); + if (!SetCSRTensorToProto(csr_tensor_abs, attr_proto)) { + return false; + } + } else if (type->isa()) { + auto coo_tensor_abs = abs->cast(); + if (!SetCOOTensorToProto(coo_tensor_abs, attr_proto)) { + return false; + } + } else { + MS_LOG(ERROR) << "Type of cnode need to be supported: " << type->type_name(); + return false; + } + return true; +} +//设置proto参数 +bool IrExportBuilder::SetAbstractToNodeProto(const CNodePtr &node, mind_ir::NodeProto *const node_proto) { + // Get shape of cnode + // 1. need to get shape from tuple element + // 2. save shape in TensorProto + MS_EXCEPTION_IF_NULL(node); + auto type = node->Type(); + auto shape = node->Shape(); + auto abs = node->abstract(); + // For the bprop fg which has not been renormalized. + if (type == nullptr || shape == nullptr) { + return true; + } + mind_ir::AttributeProto *attr_proto = node_proto->add_attribute(); + if (!SetAbstractToNodeProto(abs, attr_proto)) { + MS_LOG(ERROR) << "Set shape to NodeProto for " << node->DebugString() << " failed."; + return false; + } + attr_proto->set_name("shape"); + return true; +} +//设置proto参数 +bool IrExportBuilder::BuildCNode(const CNodePtr &node, mind_ir::GraphProto *const graph_proto) { + // 构建计算图中的一个 CNode 节点,并将其表示添加到图的 proto 中 + auto inputs_size = node->size();// 获取 CNode 的输入数量 + if (inputs_size < 1) { + MS_LOG(ERROR) << "Inputs of node " << node->DebugString() << " is empty"; + return false; + } + + // Need to build input node before dealing with cnode + // 需要先构建输入节点,然后再处理 CNode + std::vector input_names; + for (size_t i = 1; i < inputs_size; i++) { + auto input = node->input(i); + std::string node_name = BuildInputNode(input, graph_proto);// 构建输入节点并获取节点名 + if (node_name.empty()) { + MS_LOG(ERROR) << "Build input node for " << input->DebugString() << " failed."; + return false; + } + input_names.push_back(node_name);// 将输入节点名加入列表 + } + + // Build cnode + // 构建 CNode + mind_ir::NodeProto *node_proto = graph_proto->add_node();// 添加一个节点表示到图的 proto 中 + std::string output_name = GetUniqueNodeName(node);// 获取唯一的节点名 + if (nodeName_.count(output_name) > 0) { + MS_LOG(EXCEPTION) << "There is a duplicate name: " << output_name; + } + (void)nodeName_.insert(output_name);// 将节点名加入已用名字集合 + node_proto->add_output(output_name);// 设置节点的输出名 + node_proto->set_name(output_name);// 设置节点的名字 + node_proto->set_domain(node->fullname_with_scope());// 设置节点的域 + AnfNodePtr op = node->input(0);// 获取操作节点 + std::string type_name = GetOpTypeName(op);// 获取操作类型名 + if (type_name.empty()) { + MS_LOG(ERROR) << "Get op type name for " << op->DebugString() << " failed."; + return false; + } + node_proto->set_op_type(type_name);// 设置节点的操作类型 + last_node_ = node_proto;// 记录最后一个节点 + // Maybe Tensor or Function or nullptr + if (!SetAbstractToNodeProto(node, node_proto)) { + return false; + } + // 将输入节点名加入节点的输入列表中 + (void)std::for_each(input_names.begin(), input_names.end(), + [&node_proto](const string &name) { node_proto->add_input(name); }); + return true; +} + +std::string IrExportBuilder::BuildInputNode(const AnfNodePtr &node, mind_ir::GraphProto *const graph_proto) { + // Return the NodeName that the node has been processed. + auto iter = node_name_map_.find(node); + if (iter != node_name_map_.end()) { + return iter->second; + } + + std::string node_name = GetUniqueNodeName(node); + // FuncGraph will be added to functions and the input name is the function name. + if (IsValueNode(node)) { + FuncGraphPtr fg = GetValueNode(node); + todo_.push_back(fg); + return fg->ToString(); + } + if (node->isa()) { + (void)nodeName_.insert(node_name); + // When node input is a ValueNode, need to create a Constant Node + mind_ir::NodeProto *node_proto = graph_proto->add_node(); + node_proto->set_name(node_name); + node_proto->add_output(node_name); + if (!SetAttributeProto(node, node_proto)) { + return ""; + } + } + return node_name; +} +//建立输入节点 +std::string IrExportBuilder::GetUniqueNodeName(const AnfNodePtr &node) { + // Naming anfnode + // 1. parameter is unique in one func_graph + // 2. cnode and valuenode may be reduplicative, so add index to identify. + auto iter = node_name_map_.find(node); + if (iter != node_name_map_.end()) { + return iter->second; + } else { + std::string node_name = GetNodeName(node); + // Compatible before. CNode = FuncGraphName:CNodeName:index ,Parameter = FuncGraphName:ParameterName + if (node->isa()) { + node_name = node_name + ":" + std::to_string(GetUniqueID()); + } + // Avoid duplicate name. + while (nodeName_.count(node_name) > 0) { + node_name = node_name + "_" + std::to_string(GetUniqueID()); + } + node_name_map_[node] = node_name; + return node_name; + } +} +//获得未确定节点的名字 +std::string IrExportBuilder::GetNodeName(const AnfNodePtr &node) const { + MS_EXCEPTION_IF_NULL(node); + std::string node_name = ""; + if (node->func_graph() != nullptr) { + node_name = node->func_graph()->ToString() + ":"; + } + if (node->isa()) { + // Needn't value + node_name += node->AnfNode::ToString(); + } else { + node_name += node->ToString(); + } + MS_LOG(DEBUG) << "GetNodeName: " << node_name; + return node_name; +} + +bool IrExportBuilder::SetAttributeProto(const AnfNodePtr &node, mind_ir::NodeProto *const node_proto) { + if (node == nullptr || node_proto == nullptr) { + MS_LOG(EXCEPTION) << "AnfNode or NodeProto is null!"; + } + auto value_node = node->cast(); + MS_EXCEPTION_IF_NULL(value_node); + auto value = value_node->value(); + node_proto->set_op_type("Constant"); + mind_ir::AttributeProto *attr_proto = node_proto->add_attribute(); + attr_proto->set_name("value"); + MS_LOG(DEBUG) << "Set Constant attribute: " << value->ToString(); + return SetValueToAttributeProto(value, attr_proto); +} +//获得节点名 +bool IrExportBuilder::SetTypeToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) { + if (value == nullptr || attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "ValuePtr or AttributeProto is null!"; + } + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + if (value->isa()) { + tensor_proto->set_name("value0"); + auto int_value = value->cast(); + auto data_type = GetMindirDataBitsIntType(int_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + tensor_proto->set_name("value0"); + auto float_value = value->cast(); + auto data_type = GetMindirDataBitsUIntType(float_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + tensor_proto->set_name("value0"); + auto float_value = value->cast(); + auto data_type = GetMindirDataBitsFloatType(float_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + tensor_proto->set_name("value0"); + tensor_proto->set_data_type(mind_ir::TensorProto_DataType_BOOL); + } else if (value->isa()) { + tensor_proto->set_name("tensor0"); + auto elem_type = value->cast()->element(); + if (elem_type->isa()) { + auto int_value = elem_type->cast(); + auto data_type = GetMindirDataBitsIntType(int_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (elem_type->isa()) { + auto float_value = elem_type->cast(); + auto data_type = GetMindirDataBitsFloatType(float_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else { + MS_LOG(ERROR) << "Unsupported type " << elem_type->type_name(); + return false; + } + } else { + MS_LOG(EXCEPTION) << "Unsupported type: " << value->type_name(); + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetValueToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) { + if (value == nullptr || attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "ValuePtr or AttributeProto is null!"; + } + if (value->isa() || value->isa()) { + return SetScalarToAttributeProto_ir(value, attr_proto); + } else if (value->isa() || value->isa()) { + return SetTypeToAttributeProto(value, attr_proto); + } else if (value->isa()) { + if (!SetSequenceToAttributeProto(value->cast(), attr_proto)) { + MS_LOG(ERROR) << "Set sequence to AttributeProto failed."; + return false; + } + MS_LOG(DEBUG) << "Attr string: " << value->type_name(); + } else if (value->isa()) { + return SetTensorToAttributeProto(value, attr_proto); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_NONE); + MS_LOG(DEBUG) << "Attr string: " << value->type_name(); + } else if (value->isa()) { + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UMONAD); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_IOMONAD); + } else { + MS_LOG(ERROR) << "Unsupported Monad type: " << value->type_name(); + return false; + } + } else { + MS_LOG(ERROR) << "Unsupported type: " << value->type_name(); + return false; + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetScalarToAttributeProto_ir(const ValuePtr &value, + mind_ir::AttributeProto *const attr_proto) const { + if (value == nullptr || attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "ValuePtr or AttributeProto is null!"; + } + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_STRING); + attr_proto->set_s(GetValue(value)); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_BOOL); + int64_t attr_value = GetValue(value) ? 1 : 0; + attr_proto->set_i(attr_value); + } else if (SetScalarToAttributeProtoForInt_ir(value, attr_proto)) { + return true; + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_FLOAT); + attr_proto->set_f(GetValue(value)); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_DOUBLE); + attr_proto->set_d(GetValue(value)); + } else { + MS_LOG(ERROR) << "Unsupported scalar type: " << value->type_name(); + return false; + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetScalarToAttributeProtoForInt_ir(const ValuePtr &value, + mind_ir::AttributeProto *const attr_proto) const { + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT8); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT16); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT32); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT64); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT8); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT16); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT32); + attr_proto->set_i(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT64); + attr_proto->set_i(UlongToLong(value->cast()->value())); + } else { + return false; + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetTypeToAttributeProto_irs(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) { + if (attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "AttributeProto is null!"; + } + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + auto int_value = value->cast(); + auto data_type = GetMindirDataBitsIntType(int_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + auto float_value = value->cast(); + auto data_type = GetMindirDataBitsFloatType(float_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + auto uint_value = value->cast(); + auto data_type = GetMindirDataBitsFloatType(uint_value->nbits()); + if (data_type == mind_ir::TensorProto_DataType_UNDEFINED) { + return false; + } + tensor_proto->set_data_type(data_type); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + mind_ir::TensorProto *tensor_proto = attr_proto->add_tensors(); + tensor_proto->set_data_type(mind_ir::TensorProto_DataType_BOOL); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TENSORS); + return SetTensorToAttributeProto(value, attr_proto); + } else { + MS_LOG(EXCEPTION) << "Unsupported type: " << value->type_name(); + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetScalarToAttributeProto_irs(const ValuePtr &value, + mind_ir::AttributeProto *const attr_proto) const { + if (attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "AttributeProto is null!"; + } + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_STRING); + attr_proto->add_strings(GetValue(value)); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_BOOL); + attr_proto->add_ints(GetValue(value)); + } else if (SetScalarToAttributeProtoForInt_irs(value, attr_proto)) { + return true; + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_FLOAT); + attr_proto->add_floats(GetValue(value)); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_DOUBLE); + attr_proto->add_doubles(GetValue(value)); + } else { + MS_LOG(ERROR) << "Unsupported scalar type: " << value->type_name(); + return false; + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetScalarToAttributeProtoForInt_irs(const ValuePtr &value, + mind_ir::AttributeProto *const attr_proto) const { + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT8); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT16); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT32); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_INT64); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT8); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT16); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT32); + attr_proto->add_ints(value->cast()->value()); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_UINT64); + attr_proto->add_ints(SizeToInt(value->cast()->value())); + } else { + return false; + } + return true; +} +//设置proto名 +bool IrExportBuilder::SetSeqElemToAttributeProto(const ValuePtr &value, mind_ir::AttributeProto *const attr_proto) { + if (value == nullptr) { + MS_LOG(ERROR) << "Value is nullptr"; + return false; + } + if (value->isa() || value->isa()) { + return SetScalarToAttributeProto_irs(value, attr_proto); + } + return SetTypeToAttributeProto_irs(value, attr_proto); +} + +bool IrExportBuilder::SetSequenceToAttributeProto(const ValueSequencePtr &value, + mind_ir::AttributeProto *const attr_proto) { + if (value == nullptr || attr_proto == nullptr) { + MS_LOG(EXCEPTION) << "ValueSequencePtr or AttributeProto is null!"; + } + if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_TUPLE); + } else if (value->isa()) { + attr_proto->set_type(mind_ir::AttributeProto_AttributeType_LIST); + } else { + MS_LOG(EXCEPTION) << "The sequance value should be ValueTuple or ValueList, but it is " << value->ToString(); + } + auto value_sequence = value->cast(); + MS_EXCEPTION_IF_NULL(value_sequence); + const auto &values = value_sequence->value(); + if (values.empty()) { + MS_LOG(DEBUG) << "SetSequenceToAttributeProto sequence size is 0"; + return true; + } + for (const auto &item : values) { + mind_ir::AttributeProto *attr_values = attr_proto->add_values(); + MS_EXCEPTION_IF_NULL(item); + if (item->isa()) { + if (!SetSequenceToAttributeProto(item->cast(), attr_values)) { + MS_LOG(ERROR) << "Set sequence to AttributeProto failed."; + return false; + } + } else { + if (!SetSeqElemToAttributeProto(item, attr_values)) { + MS_LOG(ERROR) << "Set seq elem to AttributeProto failed."; + return false; + } + } + } + return true; +} +//设置proto名 +std::string GetBinaryProtoString(const FuncGraphPtr &func_graph) { + auto builder = std::make_shared(); + if (builder == nullptr) { + MS_LOG(ERROR) << "Create ir exporter failed!"; + return ""; + } + auto exporter = std::make_shared(builder); + if (exporter == nullptr) { + return ""; + } + auto ret = exporter->GetDumpString(func_graph); + return ret; +} +//获得protostring +bool DumpBinaryProto(const FuncGraphPtr &func_graph, const std::string &file_path, + const FuncGraphPtr ¶m_layout_fg) { + auto exporter = std::make_shared(std::make_shared()); + auto proto = exporter->GetDumpProto(func_graph, param_layout_fg); + if (proto == nullptr) { + MS_LOG(ERROR) << "Get binary proto for graph " << func_graph->ToString() << " failed."; + return false; + } + + auto realpath = Common::CreatePrefixPath(file_path, true); + if (!realpath.has_value()) { + MS_LOG(ERROR) << "Get real path of file " << file_path << " failed."; + return false; + } + + ChangeFileMode(realpath.value(), S_IWUSR); + std::ofstream fout(realpath.value()); + if (!fout.is_open()) { + MS_LOG(ERROR) << "Open the file '" << realpath.value() << "' failed!" << ErrnoToString(errno); + return false; + } + + if (!proto->SerializeToOstream(&fout)) { + MS_LOG(ERROR) << "Failed to write the mindir proto to file " << realpath.value(); + fout.close(); + return false; + } + fout.close(); + ChangeFileMode(realpath.value(), S_IRUSR); + return true; +} +} // namespace mindspore