codekingpro/portable-devtools
115k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6# This tool is not used directly in bert optimization. It could assist developing the optimization script on the following scenarios:
7# (1) It could simplify graph by removing many sub-graphs related to reshape.
8# (2) It could reduce extra inputs and outputs to fit other tools. The script compare_bert_results.py or bert_perf_test.py requires 3 inputs.
9
10import argparse
11import logging
12import os
13import re # noqa: F401
14import sys
15import tempfile
16from collections import deque # noqa: F401
17from datetime import datetime
18from pathlib import Path # noqa: F401
19
20import numpy as np
21import onnx
22from onnx import ModelProto, TensorProto, numpy_helper
23from onnx_model import OnnxModel
24
25import onnxruntime
26
27logger = logging.getLogger(__name__)
28
29CONSTANT_SHAPE_NAME_PREFIX = "constant_shape_opt__"
30RESHAPE_INPUT_SHAPE_PREFIX = "reshape_input_shape__"
31
32
33class BertOnnxModelShapeOptimizer(OnnxModel):
34 """
35 This optimizer will replace Shape output or the shape input of Reshape node by initializer. Currently, it requires
36 model inputs to have static shape.
37 """
38
39 def __init__(self, onnx_model):
40 super().__init__(onnx_model.model)
41
42 def add_shape_initializer(self, shape):
43 """
44 Add an initializer for constant shape.
45 """
46 shape_value = np.asarray(shape, dtype=np.int64)
47 constant_shape_name = self.create_node_name("Constant", CONSTANT_SHAPE_NAME_PREFIX)
48 tensor = onnx.helper.make_tensor(
49 name=constant_shape_name,
50 data_type=TensorProto.INT64,
51 dims=shape_value.shape,
52 vals=shape_value,
53 )
54 self.add_initializer(tensor)
55 return tensor
56
57 def get_shape_outputs(self):
58 """
59 Returns a list of output names of all Shape nodes.
60 """
61 input_name_to_nodes = self.input_name_to_nodes()
62
63 outputs = []
64 for node in self.model.graph.node:
65 if node.op_type == "Shape":
66 if node.output[0] in input_name_to_nodes:
67 outputs.append(node.output[0])
68
69 return outputs
70
71 def get_reshape_shape_inputs(self):
72 """
73 Returns a list of shape input names of Reshape nodes.
74 """
75 self.output_name_to_node()
76
77 shape_inputs = []
78 for node in self.model.graph.node:
79 if node.op_type == "Reshape":
80 shape_inputs.append(node.input[1])
81
82 return shape_inputs
83
84 def add_shape_for_reshape_input(self):
85 """
86 For each Reshape node, create a Shape node for its first input.
87 Returns the output names of these Shape nodes.
88 """
89 output_names = []
90 nodes_to_add = []
91 for node in self.model.graph.node:
92 if node.op_type == "Reshape":
93 input = node.input[0]
94 output_name = self.create_node_name("Reshape_Input", RESHAPE_INPUT_SHAPE_PREFIX)
95 shape_node = onnx.helper.make_node("Shape", inputs=[input], outputs=[output_name])
96 nodes_to_add.append(shape_node)
97 output_names.append(output_name)
98
99 self.add_nodes(nodes_to_add)
100 return output_names
101
102 def add_extra_graph_output(self, extra_outputs):
103 """
104 Add a list of output names to graph output.
105 """
106 names_to_evaluate = []
107 output_names = [output.name for output in self.model.graph.output]
108 for name in extra_outputs:
109 if self.get_initializer(name) is not None: # already a constant
110 continue
111 names_to_evaluate.append(name)
112
113 if name not in output_names:
114 output_info = onnx.helper.ValueInfoProto()
115 output_info.name = name
116 self.model.graph.output.extend([output_info])
117 output_names.append(name)
118
119 return names_to_evaluate
120
121 # Update input and output shape to be static
122 def use_static_input(self, inputs, batch_size=1, max_seq_len=128):
123 """
124 Update the model to use static axes instead of dynamic axes for graph inputs.
125 """
126 for input in self.model.graph.input:
127 if input.name in inputs:
128 dim_proto = input.type.tensor_type.shape.dim[0]
129 dim_proto.dim_value = batch_size
130 dim_proto = input.type.tensor_type.shape.dim[1]
131 if dim_proto.HasField("dim_param"):
132 dim_proto.dim_value = max_seq_len
133 elif dim_proto.HasField("dim_value") and dim_proto.dim_value != max_seq_len:
134 raise ValueError(
135 f"Unable to set dimension value to {max_seq_len} for axis {1} of {input.name}. Contradicts existing dimension value {dim_proto.dim_value}."
136 )
137
138 def create_dummy_inputs(
139 self,
140 input_ids,
141 segment_ids,
142 input_mask,
143 batch_size,
144 sequence_length,
145 elem_type,
146 dictionary_size=8,
147 ):
148 """
149 Create dummy data for model inputs. If the model has more than 3 inputs, please update this function accordingly before running the tool.
150 """
151 assert elem_type in [1, 6, 7] # only int32, int64 and float32 are supported.
152
153 # Create dummy inputs
154 input_1 = np.random.randint(dictionary_size, size=(batch_size, sequence_length), dtype=np.int32)
155 input_2 = np.ones((batch_size, sequence_length), dtype=np.int32)
156 input_3 = np.zeros((batch_size, sequence_length), dtype=np.int32)
157
158 # Here we assume that 3 inputs have same data type
159 if elem_type == 1: # float32
160 input_1 = np.float32(input_1)
161 input_2 = np.float32(input_2)
162 input_3 = np.float32(input_3)
163 elif elem_type == 7: # int64
164 input_1 = np.int64(input_1)
165 input_2 = np.int64(input_2)
166 input_3 = np.int64(input_3)
167
168 inputs = {input_ids: input_1, input_mask: input_2, segment_ids: input_3}
169 return inputs
170
171 def shape_optimization(
172 self,
173 temp_model_path,
174 input_ids,
175 segment_ids,
176 input_mask,
177 output_names,
178 batch_size,
179 sequence_length,
180 enable_shape_opt,
181 enable_reshape_opt,
182 verbose,
183 ):
184 self.bert_inputs = [input_ids, segment_ids, input_mask]
185
186 extra_outputs = []
187 if enable_shape_opt:
188 extra_outputs.extend(self.get_shape_outputs())
189
190 if enable_reshape_opt:
191 reshape_shape_inputs = self.get_reshape_shape_inputs()
192 reshape_input_shapes = self.add_shape_for_reshape_input()
193 extra_outputs.extend(reshape_shape_inputs)
194 extra_outputs.extend(reshape_input_shapes)
195
196 if len(extra_outputs) == 0:
197 return
198
199 names_to_evaluate = self.add_extra_graph_output(extra_outputs)
200
201 # This tool does not support dynamic axes right now.
202 self.use_static_input(self.bert_inputs, batch_size, sequence_length)
203
204 with open(temp_model_path, "wb") as out:
205 out.write(self.model.SerializeToString())
206 sess_options = onnxruntime.SessionOptions()
207 sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL
208 session = onnxruntime.InferenceSession(
209 temp_model_path,
210 sess_options,
211 providers=["CUDAExecutionProvider", "CPUExecutionProvider"],
212 )
213
214 elem_type = 7
215 for input in self.model.graph.input:
216 if input.name == input_ids:
217 elem_type = input.type.tensor_type.elem_type
218 inputs = self.create_dummy_inputs(input_ids, segment_ids, input_mask, batch_size, sequence_length, elem_type)
219
220 outputs = session.run(names_to_evaluate, inputs)
221 shapes = {}
222 for i, name in enumerate(names_to_evaluate):
223 shapes[name] = outputs[i]
224
225 logger.debug(f"shapes={shapes}")
226
227 if enable_reshape_opt:
228 for i, shape_input in enumerate(reshape_shape_inputs):
229 input_shape = reshape_input_shapes[i]
230 self.update_target_shape(shapes, shape_input, input_shape, verbose)
231
232 for name, shape in shapes.items():
233 tensor = self.add_shape_initializer(shape)
234 self.replace_input_of_all_nodes(name, tensor.name)
235
236 # Remove extra outputs, and prune all nodes not linked to output.
237 self.prune_graph(output_names)
238
239 def update_target_shape(self, shapes, shape_input, input_shape, verbose):
240 """
241 Update the target shape to use 0 to represent that dimension value does not change.
242 For example, shape of source data is (2, 5, 8) and target shape is (2, 5, 4, 2), the target shape will be updated to (0, 0, 4, 2).
243 """
244 if shape_input in shapes:
245 target_shape = shapes[shape_input]
246 else:
247 initializer = self.get_initializer(shape_input)
248 assert initializer is not None
249 target_shape = numpy_helper.to_array(initializer)
250
251 if input_shape in shapes:
252 source_shape = shapes[input_shape]
253 else:
254 initializer = self.get_initializer(input_shape)
255 assert initializer is not None
256 source_shape = numpy_helper.to_array(initializer)
257
258 new_target_shape = []
259 for i, dim_value in enumerate(target_shape):
260 if i < len(source_shape) and source_shape[i] == dim_value:
261 new_target_shape.append(0)
262 else:
263 new_target_shape.append(dim_value)
264 shapes[shape_input] = new_target_shape
265
266 logger.debug(f"source_shape={source_shape}, target_shape={target_shape}, new_target_shape={new_target_shape}")
267
268 def validate_input(self, input: str):
269 if not self.find_graph_input(input):
270 valid_names = [input.name for input in self.model.graph.input]
271 raise Exception(f"Input {input} does not exist in the graph inputs: {valid_names}")
272
273 def validate_outputs(self, output_names: list[str]):
274 valid_names = [output.name for output in self.model.graph.output]
275 for name in output_names:
276 if name not in valid_names:
277 raise Exception(f"Output {name} does not exist in the graph outputs: {valid_names}")
278
279 def optimize(
280 self,
281 output_path: str,
282 input_ids: str,
283 segment_ids: str,
284 input_mask: str,
285 enable_shape_opt: bool,
286 enable_reshape_opt: bool,
287 output_names: list[str] | None = None,
288 batch_size=1,
289 sequence_length=128,
290 verbose=False,
291 ):
292 # Skip if shape optimization has been done before.
293 for tensor in self.model.graph.initializer:
294 if tensor.name.startswith(CONSTANT_SHAPE_NAME_PREFIX):
295 logger.info("Skip shape optimization since it has been done before")
296 return
297
298 self.validate_input(input_ids)
299 self.validate_input(segment_ids)
300 self.validate_input(input_mask)
301
302 if output_names is not None:
303 self.validate_outputs(output_names)
304 self.prune_graph(output_names)
305
306 remaining_outputs = [output.name for output in self.model.graph.output]
307
308 if enable_shape_opt or enable_reshape_opt:
309 if len(self.get_graph_inputs_excluding_initializers()) != 3:
310 logger.info("Skip shape optimization since graph input number is not 3")
311 return
312
313 with tempfile.TemporaryDirectory() as temp_dir:
314 temp_file_name = "temp_{}.onnx".format(datetime.now().strftime("%m_%d-%H_%M_%S"))
315 dir = "." if verbose else temp_dir
316 temp_file = os.path.join(dir, temp_file_name)
317 self.shape_optimization(
318 temp_file,
319 input_ids,
320 segment_ids,
321 input_mask,
322 remaining_outputs,
323 batch_size,
324 sequence_length,
325 enable_shape_opt,
326 enable_reshape_opt,
327 verbose,
328 )
329 logger.debug(f"Temp model with additional outputs: {temp_file}")
330 logger.warning(
331 f"Shape optimization is done. The optimized model might only work for input with batch_size={batch_size} sequence_length={sequence_length}"
332 )
333
334 if output_path is not None:
335 with open(output_path, "wb") as out:
336 out.write(self.model.SerializeToString())
337
338
339def parse_arguments():
340 parser = argparse.ArgumentParser()
341 parser.add_argument("--input", required=True, type=str)
342 parser.add_argument("--output", required=True, type=str)
343 parser.add_argument("--input_ids", required=True, type=str)
344 parser.add_argument("--segment_ids", required=True, type=str)
345 parser.add_argument("--input_mask", required=True, type=str)
346 parser.add_argument("--output_names", required=False, type=str, default=None)
347 parser.add_argument("--batch_size", required=False, type=int, default=1)
348 parser.add_argument("--sequence_length", required=False, type=int, default=128)
349 parser.add_argument("--enable_shape_opt", required=False, action="store_true")
350 parser.set_defaults(enable_shape_opt=False)
351 parser.add_argument("--enable_reshape_opt", required=False, action="store_true")
352 parser.set_defaults(enable_reshape_opt=False)
353 parser.add_argument("--verbose", required=False, action="store_true")
354 parser.set_defaults(verbose=False)
355 args = parser.parse_args()
356 return args
357
358
359def setup_logging(verbose):
360 log_handler = logging.StreamHandler(sys.stdout)
361 if verbose:
362 log_handler.setFormatter(logging.Formatter("[%(filename)s:%(lineno)s - %(funcName)20s()] %(message)s"))
363 logging_level = logging.DEBUG
364 else:
365 log_handler.setFormatter(logging.Formatter("%(filename)20s: %(message)s"))
366 logging_level = logging.INFO
367 log_handler.setLevel(logging_level)
368 logger.addHandler(log_handler)
369 logger.setLevel(logging_level)
370
371
372def main():
373 args = parse_arguments()
374 setup_logging(args.verbose)
375
376 output_names = None if args.output_names is None else args.output_names.split(";")
377
378 model = ModelProto()
379 with open(args.input, "rb") as input_file:
380 model.ParseFromString(input_file.read())
381 onnx_model = OnnxModel(model)
382
383 optimizer = BertOnnxModelShapeOptimizer(onnx_model)
384
385 optimizer.optimize(
386 args.output,
387 args.input_ids,
388 args.segment_ids,
389 args.input_mask,
390 args.enable_shape_opt,
391 args.enable_reshape_opt,
392 output_names,
393 args.batch_size,
394 args.sequence_length,
395 args.verbose,
396 )
397
398
399if __name__ == "__main__":
400 main()
401 