mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
Topo sort the model before saving (#7913)
* checkin toposort * review comments * revert and add TODO
This commit is contained in:
parent
c45ac166d3
commit
4f82ad1b58
2 changed files with 63 additions and 1 deletions
|
|
@ -112,7 +112,7 @@ class ONNXModel:
|
|||
|
||||
def find_node_by_name(self, node_name, new_nodes_list, graph):
|
||||
'''
|
||||
Find out if a node exists in a graph or a node is in the
|
||||
Find out if a node exists in a graph or a node is in the
|
||||
new set of nodes created during quantization. Return the node found.
|
||||
'''
|
||||
graph_nodes_list = list(graph.node) #deep copy
|
||||
|
|
@ -256,6 +256,8 @@ class ONNXModel:
|
|||
return True
|
||||
return False
|
||||
|
||||
# TODO:use OnnxModel.graph_topological_sort(self.model.graph) from transformers.onnx_model
|
||||
# Currently it breaks Openvino/Linux training gpu pipeline so hold off for 1.8 release
|
||||
def topological_sort(self):
|
||||
deps_count = [0]*len(self.nodes()) # dependency count of each node
|
||||
deps_to_nodes = {} # input to node indice
|
||||
|
|
|
|||
|
|
@ -781,7 +781,67 @@ class OnnxModel:
|
|||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def graph_topological_sort(graph):
|
||||
deps_count = [0]*len(graph.node) # dependency count of each node
|
||||
deps_to_nodes = {} # input to node indice
|
||||
sorted_nodes = [] # initialize sorted_nodes
|
||||
for node_idx, node in enumerate(graph.node):
|
||||
# CANNOT use len(node.input) directly because input can be optional
|
||||
deps_count[node_idx] = sum(1 for _ in node.input if _ )
|
||||
if deps_count[node_idx] == 0: # Constant doesn't depend on any inputs
|
||||
sorted_nodes.append(graph.node[node_idx])
|
||||
continue
|
||||
|
||||
for input_name in node.input:
|
||||
if input_name not in deps_to_nodes:
|
||||
deps_to_nodes[input_name] = [node_idx]
|
||||
else:
|
||||
deps_to_nodes[input_name].append(node_idx)
|
||||
|
||||
initializer_names = [init.name for init in graph.initializer]
|
||||
graph_input_names = [input.name for input in graph.input]
|
||||
input_names = initializer_names + graph_input_names
|
||||
input_names.sort()
|
||||
prev_input_name = None
|
||||
for input_name in input_names:
|
||||
if prev_input_name == input_name:
|
||||
continue
|
||||
|
||||
prev_input_name = input_name
|
||||
if input_name in deps_to_nodes:
|
||||
for node_idx in deps_to_nodes[input_name]:
|
||||
deps_count[node_idx] = deps_count[node_idx] - 1
|
||||
if deps_count[node_idx] == 0:
|
||||
sorted_nodes.append(graph.node[node_idx])
|
||||
|
||||
start = 0
|
||||
end = len(sorted_nodes)
|
||||
|
||||
while start < end:
|
||||
for output in sorted_nodes[start].output:
|
||||
if output in deps_to_nodes:
|
||||
for node_idx in deps_to_nodes[output]:
|
||||
deps_count[node_idx] = deps_count[node_idx] - 1
|
||||
if deps_count[node_idx] == 0:
|
||||
sorted_nodes.append(graph.node[node_idx])
|
||||
end = end + 1
|
||||
start = start + 1
|
||||
|
||||
assert(end == len(graph.node)), "Graph is not a DAG"
|
||||
graph.ClearField('node')
|
||||
graph.node.extend(sorted_nodes)
|
||||
|
||||
def topological_sort(self):
|
||||
#TODO: support graph_topological_sort() in subgraphs
|
||||
#for graph in self.graphs():
|
||||
# self.graph_topological_sort(graph)
|
||||
OnnxModel.graph_topological_sort(self.model.graph)
|
||||
|
||||
def save_model_to_file(self, output_path, use_external_data_format=False):
|
||||
logger.info(f"Sort graphs in topological order")
|
||||
self.topological_sort()
|
||||
|
||||
logger.info(f"Output model to {output_path}")
|
||||
|
||||
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
|
|
|
|||
Loading…
Reference in a new issue