codekingpro/portable-devtools
114k
1# --------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from pathlib import Path
6
7import onnx
8import onnx.helper as onnx_helper
9import onnx.numpy_helper as onnx_numpy_helper
10from onnx.onnx_pb import ModelProto
11
12from .quant_utils import attribute_to_kwarg, find_by_name
13
14
15def _clean_initializers_helper(graph, model):
16 """Clean unused initializers from graph.
17
18 Returns:
19 A cleaned graph without unused initializers
20 A list of tensor names, which are not produced by this graph and its subgraphes
21 """
22 requesting_tensor_names = set()
23 requesting_tensor_names.update(input_name for node in graph.node for input_name in node.input if input_name)
24 requesting_tensor_names.update(g_out.name for g_out in graph.output if g_out.name)
25
26 new_nodes = []
27 for node in graph.node:
28 new_node = node
29 graph_attrs = [
30 attr
31 for attr in node.attribute
32 if attr.type == onnx.AttributeProto.GRAPH or attr.type == onnx.AttributeProto.GRAPHS
33 ]
34 if graph_attrs:
35 kwargs = {}
36 for attr in node.attribute:
37 new_attribute = {}
38 if attr.type == onnx.AttributeProto.GRAPH:
39 (
40 cleaned_sub_graph,
41 sub_requesting_tensor_names,
42 ) = _clean_initializers_helper(attr.g, model)
43 new_attribute = {attr.name: cleaned_sub_graph}
44 requesting_tensor_names.update(sub_requesting_tensor_names)
45 elif attr.type == onnx.AttributeProto.GRAPHS:
46 cleaned_graphes = []
47 for subgraph in attr.graphs:
48 (
49 cleaned_sub_graph,
50 sub_requesting_tensor_names,
51 ) = _clean_initializers_helper(subgraph, model)
52 cleaned_graphes.append(cleaned_sub_graph)
53 requesting_tensor_names.update(sub_requesting_tensor_names)
54 new_attribute = {attr.name: cleaned_graphes}
55 else:
56 new_attribute = attribute_to_kwarg(attr)
57 kwargs.update(new_attribute)
58 new_node = onnx_helper.make_node(node.op_type, node.input, node.output, name=node.name, **kwargs)
59 new_nodes.append(new_node)
60
61 graph.ClearField("node")
62 graph.node.extend(new_nodes)
63
64 requesting_tensor_names.difference_update(output for node in graph.node for output in node.output)
65
66 unused_initializer = []
67 for initializer in graph.initializer:
68 if initializer.name in requesting_tensor_names:
69 requesting_tensor_names.remove(initializer.name)
70 else:
71 # mark it to remove, remove here directly will cause mis-behavier
72 unused_initializer.append(initializer)
73
74 name_to_input = {input.name: input for input in graph.input}
75 for initializer in unused_initializer:
76 graph.initializer.remove(initializer)
77 if initializer.name in name_to_input:
78 try:
79 graph.input.remove(name_to_input[initializer.name])
80 except StopIteration:
81 if model.ir_version < 4:
82 print(f"Warning: invalid weight name {initializer.name} found in the graph (not a graph input)")
83
84 requesting_tensor_names.difference_update(input.name for input in graph.input)
85
86 return graph, requesting_tensor_names
87
88
89class ONNXModel:
90 def __init__(self, model: ModelProto):
91 self.model = model
92
93 def nodes(self):
94 return self.model.graph.node
95
96 def initializer(self):
97 return self.model.graph.initializer
98
99 def initializer_extend(self, inits):
100 if len(inits) == 0:
101 raise ValueError("Can add an empty list.")
102 for init in self.initializer():
103 self._check_init(init, "gain")
104 for init in inits:
105 self._check_init(init)
106 self.model.graph.initializer.append(init)
107
108 def graph(self):
109 return self.model.graph
110
111 def ir_version(self):
112 return self.model.ir_version
113
114 def opset_import(self):
115 return self.model.opset_import
116
117 def set_opset_import(self, domain, version):
118 for opset in self.model.opset_import:
119 if opset.domain == domain:
120 opset.version = version
121 return
122
123 self.model.opset_import.extend([onnx_helper.make_opsetid(domain, version)])
124
125 def remove_node(self, node):
126 if node in self.model.graph.node:
127 self.model.graph.node.remove(node)
128
129 def remove_nodes(self, nodes_to_remove):
130 for node in nodes_to_remove:
131 self.remove_node(node)
132
133 def add_node(self, node):
134 self.model.graph.node.extend([self._check_node(node)])
135
136 def add_nodes(self, nodes_to_add):
137 for node in nodes_to_add:
138 self.add_node(node)
139
140 def add_initializer(self, tensor):
141 if find_by_name(tensor.name, self.model.graph.initializer) is None:
142 self._check_init(tensor)
143 self.model.graph.initializer.extend([tensor])
144
145 def get_initializer(self, name):
146 for tensor in self.model.graph.initializer:
147 if tensor.name == name:
148 return tensor
149 return None
150
151 def find_graph_input(self, input_name):
152 for input in self.model.graph.input:
153 if input.name == input_name:
154 return input
155 return None
156
157 def find_graph_output(self, output_name):
158 for output in self.model.graph.output:
159 if output.name == output_name:
160 return output
161 return None
162
163 def get_tensor_type(self, tensor_name: str):
164 tensor_type_map = {obj.name: obj.type for obj in self.model.graph.value_info}
165
166 if tensor_name in tensor_type_map:
167 return tensor_type_map[tensor_name].tensor_type
168
169 g_input = self.find_graph_input(tensor_name)
170 if g_input:
171 return g_input.type.tensor_type
172
173 g_output = self.find_graph_output(tensor_name)
174 if g_output:
175 return g_output.type.tensor_type
176
177 return None
178
179 def get_constant_value(self, output_name):
180 for node in self.model.graph.node:
181 if node.op_type == "Constant":
182 if node.output[0] == output_name:
183 for attr in node.attribute:
184 if attr.name == "value":
185 return onnx_numpy_helper.to_array(attr.t)
186
187 # Fallback to initializer since constant folding may have been applied.
188 initializer = self.get_initializer(output_name)
189 if initializer is not None:
190 return onnx_numpy_helper.to_array(initializer)
191
192 return None
193
194 def get_initializer_name_set(self):
195 return {initializer.name for initializer in self.model.graph.initializer}
196
197 def remove_initializer(self, tensor):
198 if tensor in self.model.graph.initializer:
199 self.model.graph.initializer.remove(tensor)
200 for input in self.model.graph.input:
201 if input.name == tensor.name:
202 self.model.graph.input.remove(input)
203 break
204
205 def remove_initializers(self, init_to_remove):
206 for initializer in init_to_remove:
207 self.remove_initializer(initializer)
208
209 def get_non_initializer_inputs(self):
210 initializer_names = self.get_initializer_name_set()
211 non_initializer_inputs = set()
212 for input in self.model.graph.input:
213 if input.name not in initializer_names:
214 non_initializer_inputs.add(input.name)
215 return non_initializer_inputs
216
217 def input_name_to_nodes(self):
218 input_name_to_nodes = {}
219 for node in self.model.graph.node:
220 for input_name in node.input:
221 if input_name: # Could be empty when it is optional
222 if input_name not in input_name_to_nodes:
223 input_name_to_nodes[input_name] = [node]
224 else:
225 input_name_to_nodes[input_name].append(node)
226 return input_name_to_nodes
227
228 def output_name_to_node(self):
229 output_name_to_node = {}
230 for node in self.model.graph.node:
231 for output_name in node.output:
232 if output_name: # Could be empty when it is optional
233 output_name_to_node[output_name] = node
234 return output_name_to_node
235
236 def get_children(self, node, input_name_to_nodes=None):
237 if input_name_to_nodes is None:
238 input_name_to_nodes = self.input_name_to_nodes()
239
240 children = []
241 for output in node.output:
242 if output in input_name_to_nodes:
243 for node in input_name_to_nodes[output]:
244 children.append(node) # noqa: PERF402
245 return children
246
247 def get_parents(self, node, output_name_to_node=None):
248 if output_name_to_node is None:
249 output_name_to_node = self.output_name_to_node()
250
251 parents = []
252 for input in node.input:
253 if input in output_name_to_node:
254 parents.append(output_name_to_node[input])
255 return parents
256
257 def get_parent(self, node, idx, output_name_to_node=None):
258 if output_name_to_node is None:
259 output_name_to_node = self.output_name_to_node()
260
261 if len(node.input) <= idx:
262 return None
263
264 input = node.input[idx]
265 if input not in output_name_to_node:
266 return None
267
268 return output_name_to_node[input]
269
270 def find_node_by_name(self, node_name, new_nodes_list, graph):
271 """Find out if a node exists in a graph or a node is in the
272 new set of nodes created during quantization.
273
274 Returns:
275 The node found or None.
276 """
277 graph_nodes_list = list(graph.node) # deep copy
278 graph_nodes_list.extend(new_nodes_list)
279 node = find_by_name(node_name, graph_nodes_list)
280 return node
281
282 def get_largest_node_name_suffix(self, node_name_prefix):
283 """
284 Gets the largest node name (int) suffix for all node names that begin with `node_name_prefix`.
285 Example: for nodes my_prefix_0 and my_prefix_3, this method returns 3.
286 """
287 suffix = -1
288
289 for node in self.model.graph.node:
290 if node.name and node.name.startswith(node_name_prefix):
291 try:
292 index = int(node.name[len(node_name_prefix) :])
293 suffix = max(index, suffix)
294 except ValueError:
295 continue
296
297 return suffix
298
299 def get_largest_initializer_name_suffix(self, initializer_name_prefix):
300 """
301 Gets the largest initializer name integer suffix for all initializer names that begin
302 with `initializer_name_prefix`. This can be used to create unique initializer names.
303
304 Example: for initializer names 'my_weight_0' and 'my_weight_3', this method returns 3 if
305 `initializer_name_prefix` is 'my_weight_'.
306 """
307 suffix = -1
308
309 for initializer in self.model.graph.initializer:
310 if initializer.name.startswith(initializer_name_prefix):
311 try:
312 index = int(initializer.name[len(initializer_name_prefix) :])
313 suffix = max(index, suffix)
314 except ValueError:
315 continue
316
317 return suffix
318
319 def find_nodes_by_initializer(self, graph, initializer):
320 """
321 Find all nodes with given initializer as an input.
322 """
323 nodes = []
324 for node in graph.node:
325 for node_input in node.input:
326 if node_input == initializer.name:
327 nodes.append(node)
328 return nodes
329
330 @staticmethod
331 def __get_initializer(name, graph_path):
332 for gid in range(len(graph_path) - 1, -1, -1):
333 graph = graph_path[gid]
334 for tensor in graph.initializer:
335 if tensor.name == name:
336 return tensor, graph
337 return None, None
338
339 @staticmethod
340 def __replace_gemm_with_matmul(graph_path):
341 new_nodes = []
342 graph = graph_path[-1]
343 for node in graph.node:
344 graph_attrs = [attr for attr in node.attribute if attr.type == 5 or attr.type == 10]
345 if graph_attrs:
346 kwargs = {}
347 for attr in node.attribute:
348 if attr.type == 5:
349 graph_path.append(attr.g)
350 kv = {attr.name: ONNXModel.__replace_gemm_with_matmul(graph_path)}
351 elif attr.type == 10:
352 value = []
353 for subgraph in attr.graphs:
354 graph_path.append(subgraph)
355 value.extend([ONNXModel.__replace_gemm_with_matmul(graph_path)])
356 kv = {attr.name: value}
357 else:
358 kv = attribute_to_kwarg(attr)
359 kwargs.update(kv)
360 node = onnx_helper.make_node( # noqa: PLW2901
361 node.op_type, node.input, node.output, name=node.name, **kwargs
362 )
363
364 if node.op_type == "Gemm":
365 alpha = 1.0
366 beta = 1.0
367 transA = 0 # noqa: N806
368 transB = 0 # noqa: N806
369 for attr in node.attribute:
370 if attr.name == "alpha":
371 alpha = onnx_helper.get_attribute_value(attr)
372 elif attr.name == "beta":
373 beta = onnx_helper.get_attribute_value(attr)
374 elif attr.name == "transA":
375 transA = onnx_helper.get_attribute_value(attr) # noqa: N806
376 elif attr.name == "transB":
377 transB = onnx_helper.get_attribute_value(attr) # noqa: N806
378 if alpha == 1.0 and beta == 1.0 and transA == 0:
379 inputB = node.input[1] # noqa: N806
380 if transB == 1:
381 B, Bs_graph = ONNXModel.__get_initializer(node.input[1], graph_path) # noqa: N806
382 if B:
383 # assume B is not used by any other node
384 B_array = onnx_numpy_helper.to_array(B) # noqa: N806
385 B_trans = onnx_numpy_helper.from_array(B_array.T) # noqa: N806
386 B_trans.name = B.name
387 Bs_graph.initializer.remove(B)
388 for input in Bs_graph.input:
389 if input.name == inputB:
390 Bs_graph.input.remove(input)
391 break
392 Bs_graph.initializer.extend([B_trans])
393 else:
394 inputB += "_Transposed" # noqa: N806
395 transpose_node = onnx_helper.make_node(
396 "Transpose",
397 inputs=[node.input[1]],
398 outputs=[inputB],
399 name=node.name + "_Transpose" if node.name else "",
400 )
401 new_nodes.append(transpose_node)
402
403 matmul_node = onnx_helper.make_node(
404 "MatMul",
405 inputs=[node.input[0], inputB],
406 outputs=[node.output[0] + ("_MatMul" if len(node.input) > 2 else "")],
407 name=node.name + "_MatMul" if node.name else "",
408 )
409 new_nodes.append(matmul_node)
410
411 if len(node.input) > 2:
412 add_node = onnx_helper.make_node(
413 "Add",
414 inputs=[node.output[0] + "_MatMul", node.input[2]],
415 outputs=node.output,
416 name=node.name + "_Add" if node.name else "",
417 )
418 new_nodes.append(add_node)
419
420 # unsupported
421 else:
422 new_nodes.append(node)
423
424 # not GEMM
425 else:
426 new_nodes.append(node)
427
428 graph.ClearField("node")
429 graph.node.extend(new_nodes)
430 graph_path.pop()
431 return graph
432
433 def replace_gemm_with_matmul(self):
434 graph_path = [self.graph()]
435 ONNXModel.__replace_gemm_with_matmul(graph_path)
436
437 def save_model_to_file(self, output_path, use_external_data_format=False):
438 """
439 Save model to external data, which is needed for model size > 2GB
440 """
441 self.topological_sort()
442 if use_external_data_format:
443 onnx.external_data_helper.convert_model_to_external_data(
444 self.model,
445 all_tensors_to_one_file=True,
446 location=Path(output_path).name + ".data",
447 convert_attribute=True,
448 )
449 for init in self.model.graph.initializer:
450 self._check_init(init, "end")
451 onnx.save_model(self.model, output_path)
452
453 @staticmethod
454 def replace_node_input(node, old_input_name, new_input_name):
455 assert isinstance(old_input_name, str) and isinstance(new_input_name, str)
456 for j in range(len(node.input)):
457 if node.input[j] == old_input_name:
458 node.input[j] = new_input_name
459
460 def replace_input_of_all_nodes(self, old_input_name, new_input_name):
461 for node in self.model.graph.node:
462 ONNXModel.replace_node_input(node, old_input_name, new_input_name)
463
464 def replace_input_of_nodes(self, old_input_name, new_input_name, node_names_set):
465 for node in self.model.graph.node:
466 if node.name in node_names_set:
467 ONNXModel.replace_node_input(node, old_input_name, new_input_name)
468
469 @staticmethod
470 def replace_node_output(node, old_output_name, new_output_name):
471 assert isinstance(old_output_name, str) and isinstance(new_output_name, str)
472 for j in range(len(node.output)):
473 if node.output[j] == old_output_name:
474 node.output[j] = new_output_name
475
476 def replace_output_of_all_nodes(self, old_output_name, new_output_name):
477 for node in self.model.graph.node:
478 ONNXModel.replace_node_output(node, old_output_name, new_output_name)
479
480 def replace_output_of_nodes(self, old_output_name, new_output_name, node_names_set):
481 for node in self.model.graph.node:
482 if node.name in node_names_set:
483 ONNXModel.replace_node_output(node, old_output_name, new_output_name)
484
485 def remove_unused_constant(self):
486 input_name_to_nodes = self.input_name_to_nodes()
487
488 # remove unused constant
489 unused_nodes = []
490 nodes = self.nodes()
491 for node in nodes:
492 if (
493 node.op_type == "Constant"
494 and not self.is_graph_output(node.output[0])
495 and node.output[0] not in input_name_to_nodes
496 ):
497 unused_nodes.append(node)
498
499 self.remove_nodes(unused_nodes)
500
501 ununsed_weights = []
502 for w in self.initializer():
503 if w.name not in input_name_to_nodes and not self.is_graph_output(w.name):
504 ununsed_weights.append(w)
505 # Remove from graph.input
506 for graph_input in self.graph().input:
507 if graph_input.name == w.name:
508 self.graph().input.remove(graph_input)
509
510 self.remove_initializers(ununsed_weights)
511
512 def is_graph_output(self, output_name):
513 return any(output.name == output_name for output in self.model.graph.output)
514
515 def is_graph_input(self, tensor_name: str) -> bool:
516 return any(input.name == tensor_name for input in self.model.graph.input)
517
518 # TODO:use OnnxModel.graph_topological_sort(self.model.graph) from transformers.onnx_model
519 # Currently it breaks Openvino/Linux training gpu pipeline so hold off for 1.8 release
520 def topological_sort(self):
521 deps_count = [0] * len(self.nodes()) # dependency count of each node
522 deps_to_nodes = {} # input to node indice
523 sorted_nodes = [] # initialize sorted_nodes
524 for node_idx, node in enumerate(self.nodes()):
525 # CANNOT use len(node.input) directly because input can be optional
526 deps_count[node_idx] = sum(1 for _ in node.input if _)
527 if deps_count[node_idx] == 0: # Constant doesn't depend on any inputs
528 sorted_nodes.append(self.nodes()[node_idx])
529 continue
530
531 for input_name in node.input:
532 if not input_name:
533 continue
534 if input_name not in deps_to_nodes:
535 deps_to_nodes[input_name] = [node_idx]
536 else:
537 deps_to_nodes[input_name].append(node_idx)
538
539 initializer_names = [init.name for init in self.initializer()]
540 graph_input_names = [input.name for input in self.model.graph.input]
541 input_names = initializer_names + graph_input_names
542 input_names.sort()
543 prev_input_name = None
544 for input_name in input_names:
545 if prev_input_name == input_name:
546 continue
547
548 prev_input_name = input_name
549 if input_name in deps_to_nodes:
550 for node_idx in deps_to_nodes[input_name]:
551 deps_count[node_idx] = deps_count[node_idx] - 1
552 if deps_count[node_idx] == 0:
553 sorted_nodes.append(self.nodes()[node_idx])
554
555 start = 0
556 end = len(sorted_nodes)
557
558 while start < end:
559 for output in sorted_nodes[start].output:
560 if output in deps_to_nodes:
561 for node_idx in deps_to_nodes[output]:
562 deps_count[node_idx] = deps_count[node_idx] - 1
563 if deps_count[node_idx] == 0:
564 sorted_nodes.append(self.nodes()[node_idx])
565 end = end + 1
566 start = start + 1
567
568 assert end == len(self.graph().node), "Graph is not a DAG"
569 self.graph().ClearField("node")
570 self.graph().node.extend(sorted_nodes)
571
572 def clean_initializers(self):
573 return _clean_initializers_helper(self.graph(), self.model)
574
575 def _check_init(self, init, test=None):
576 if init.data_type == onnx.TensorProto.FLOAT8E4M3FN:
577 if init.HasField("raw_data"):
578 b = list(init.raw_data)
579 if any((i & 127) == 127 for i in b):
580 raise ValueError(f"Initializer {init.name!r} has nan.")
581 return init
582
583 def _check_node(self, node):
584 """
585 A quantization to float 8 does not use quantized bias but float 16 bias.
586 This function checks that DequantizeLinear is not used to
587 dequantize from float 16.
588 """
589 if node.op_type == "DequantizeLinear":
590 zero_point = node.input[2]
591 init = self.get_initializer(zero_point)
592 dtype = init.data_type
593 if dtype in {
594 onnx.TensorProto.FLOAT16,
595 onnx.TensorProto.FLOAT,
596 onnx.TensorProto.DOUBLE,
597 onnx.TensorProto.BFLOAT16,
598 }:
599 raise RuntimeError(f"Unsupported DequantizeLinear operator, dequantization from {dtype}.")
600 return node
601 