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()方法来实现整个读取节点的方式。