Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
shape_optimizer.py401 linesDownload Raw Back to transformers
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 
codekingpro/portable-devtools · Team Ai