如何读取MindSpore中的.pb文件中的节点

1.系统环境

硬件环境(Ascend/GPU/CPU): GPU
MindSpore版本: mindspore=2.0
执行模式(PyNative/ Graph):不限
Python版本:3.7
操作系统平台:Linux

2. 解决方案

修改逻辑:首先,需要适配节点的node_mapper匹配算法的逻辑,则需要以下Node的参数。

class Node:
    """
    Computational graph node, a node represents an operator.

    Args:
        name (str): Node name.
        node_id (str): Node(operator) id.
        node_type (str): Node(operator) type.
        scope (str, optional): Scope.
    """

    def __init__(self, name, node_id, node_type, scope=None):
        self.name = name
        self.node_id = node_id
        self.node_type = node_type
        self.scope = scope
        self.shape = None
        self.label = None
        self.from_nodes = set()
        self.to_nodes = set()
        self.visited = False
        self.depth = -1
        self.homo_depth = -1
        self.internal_id = -1

        self.match_idx = -1
        self.match_sim = 0
        self.tmp_match_sim = 0
        self.tmp_matched_node = None
        self.partition = None
        self.footprints = None

其次,原本.pb文件中node的name属性即是各个node的id,full_name属性才是node的name,而.pbtxt文件中,没有node_id,name属性就是该节点的名称。同时,发现topo_id实际上就是直接取的node_id的值,除了计数外并没有其他作用,故直接设置一个新的id参数,每过一个节点迭代+1,仅用作节点计数区分使用。

之后,由于.pbtxt文件中没有parameters参数的节点,删除相应的处理操作,并且,设置新的node_map参数,通过内部存储各个节点的name来找出各个节点之间边的关系。

最终,通过onnx的ModelProto()方法来实现整个读取节点的方式。