codekingpro/portable-devtools
115k
1# --------------------------------------------------------------------------
2# Copyright (c) Microsoft, Intel Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7
8import logging
9import tempfile
10import traceback
11from pathlib import Path
12
13import onnx
14
15import onnxruntime
16from onnxruntime.tools.symbolic_shape_infer import SymbolicShapeInference
17from onnxruntime.transformers.onnx_utils import extract_raw_data_from_model, has_external_data
18
19from .fusions import ReplaceUpsampleWithResize
20from .onnx_model import ONNXModel
21from .quant_utils import add_pre_process_metadata, save_and_reload_model_with_shape_infer
22
23logger = logging.getLogger(__name__)
24
25
26def quant_pre_process(
27 input_model: str | Path | onnx.ModelProto | None = None,
28 output_model_path: str | Path | None = None,
29 skip_optimization: bool = False,
30 skip_onnx_shape: bool = False,
31 skip_symbolic_shape: bool = False,
32 auto_merge: bool = False,
33 int_max: int = 2**31 - 1,
34 guess_output_rank: bool = False,
35 verbose: int = 0,
36 save_as_external_data: bool = False,
37 all_tensors_to_one_file: bool = False,
38 external_data_location: str | None = None,
39 external_data_size_threshold: int = 1024,
40 **deprecated_kwargs,
41) -> None:
42 """Shape inference and model optimization, in preparation for quantization.
43
44 Args:
45 input_model: Path to the input model file or ModelProto
46 output_model_path: Path to the output model file
47 skip_optimization: Skip model optimization step if true. This may result in ONNX shape
48 inference failure for some models.
49 skip_onnx_shape: Skip ONNX shape inference. Symbolic shape inference is most effective
50 with transformer based models. Skipping all shape inferences may
51 reduce the effectiveness of quantization, as a tensor with unknown
52 shape can not be quantized.
53 skip_symbolic_shape: Skip symbolic shape inference. Symbolic shape inference is most
54 effective with transformer based models. Skipping all shape
55 inferences may reduce the effectiveness of quantization, as a tensor
56 with unknown shape can not be quantized.
57 auto_merge: For symbolic shape inference, automatically merge symbolic dims when
58 conflict happens.
59 int_max: For symbolic shape inference, specify the maximum value for integer to be
60 treated as boundless for ops like slice
61 guess_output_rank: Guess output rank to be the same as input 0 for unknown ops
62 verbose: Logs detailed info of inference, 0: turn off, 1: warnings, 3: detailed
63 save_as_external_data: Saving an ONNX model to external data
64 all_tensors_to_one_file: Saving all the external data to one file
65 external_data_location: The file location to save the external file
66 external_data_size_threshold: The size threshold for external data
67 """
68
69 if input_model is None:
70 input_model = deprecated_kwargs.pop("input_model_path", None)
71 assert input_model is not None
72
73 assert output_model_path is not None, "output_model_path is required."
74
75 with tempfile.TemporaryDirectory(prefix="pre.quant.") as quant_tmp_dir:
76 temp_path = Path(quant_tmp_dir)
77 model = input_model if isinstance(input_model, onnx.ModelProto) else onnx.load(input_model)
78
79 # Since Upsample is deprecated after opset v10, and the model's opset will
80 # be upgraded to at least v11 during quantization, we need to replace Upsample
81 # with Resize first to avoid generating an invalid model.
82 ai_onnx_domain = [opset for opset in model.opset_import if not opset.domain or opset.domain == "ai.onnx"]
83 if len(ai_onnx_domain) == 1:
84 opset_version = ai_onnx_domain[0].version
85 if opset_version <= 10:
86 ReplaceUpsampleWithResize(ONNXModel(model), opset_version).apply()
87 model = onnx.version_converter.convert_version(model, 11)
88 model = save_and_reload_model_with_shape_infer(model)
89
90 if not skip_symbolic_shape:
91 logger.info("Performing symbolic shape inference...")
92 model = SymbolicShapeInference.infer_shapes(
93 model,
94 int_max,
95 auto_merge,
96 guess_output_rank,
97 verbose,
98 )
99
100 if not skip_optimization:
101 # Use ORT optimizers (native code) to optimize model
102 if not skip_symbolic_shape:
103 # Need to save the inferenced model to file so as to run the optimizer
104 input_model = str(temp_path / "symbolic_shape_inferred.onnx")
105 if save_as_external_data:
106 onnx.save_model(
107 model,
108 input_model,
109 save_as_external_data=True,
110 all_tensors_to_one_file=all_tensors_to_one_file,
111 size_threshold=external_data_size_threshold,
112 convert_attribute=False,
113 )
114 else:
115 onnx.save(model, input_model)
116 model = None
117
118 opt_model_path = str(temp_path / "optimized.onnx")
119 try:
120 sess_option = onnxruntime.SessionOptions()
121 sess_option.optimized_model_filepath = opt_model_path
122 sess_option.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC
123 # For large model, extract external data from model and add to session options
124 if isinstance(input_model, onnx.ModelProto):
125 if has_external_data(input_model):
126 raise ValueError(
127 "ModelProto has external data not loaded into memory, ORT cannot create session. "
128 "Please load external data before calling this function. "
129 "See https://onnx.ai/onnx/repo-docs/ExternalData.html for more information."
130 )
131 external_names, external_values = extract_raw_data_from_model(input_model)
132 sess_option.add_external_initializers(list(external_names), list(external_values))
133 input_model = input_model.SerializeToString()
134 # the saved optimized model otherwise points to the original external data file name
135 # which is not available relative to the optimized model file
136 elif skip_symbolic_shape and save_as_external_data:
137 sess_option.add_session_config_entry(
138 "session.optimized_model_external_initializers_file_name", "optimized.onnx.data"
139 )
140
141 sess = onnxruntime.InferenceSession(input_model, sess_option, providers=["CPUExecutionProvider"])
142 # Close the session to avoid the cleanup error on Windows for temp folders
143 # https://github.com/microsoft/onnxruntime/issues/17627
144 del sess
145 except Exception:
146 logger.error(
147 "ONNX Runtime Model Optimization Failed! Consider rerun with option `--skip_optimization'."
148 )
149 logger.error(traceback.format_exc())
150
151 input_model = opt_model_path
152
153 if not skip_onnx_shape:
154 # ONNX shape inference.
155 # According to docs, infer_shapes_path should be used for 2G+ models.
156 # If the skip optimization is specified, we could be dealing with a
157 # large model. So be on the safe side, save the model
158 if model is not None:
159 input_model = str(temp_path / "symbolic_shape_inferred.onnx")
160 if save_as_external_data:
161 onnx.save_model(
162 model,
163 input_model,
164 save_as_external_data=True,
165 all_tensors_to_one_file=all_tensors_to_one_file,
166 size_threshold=external_data_size_threshold,
167 convert_attribute=False,
168 )
169 else:
170 onnx.save(model, input_model)
171 model = None
172
173 if isinstance(input_model, onnx.ModelProto):
174 input_model = str(Path(quant_tmp_dir) / "model_input.onnx")
175 onnx.save_model(
176 model,
177 input_model,
178 save_as_external_data=True,
179 all_tensors_to_one_file=all_tensors_to_one_file,
180 size_threshold=external_data_size_threshold,
181 convert_attribute=False,
182 )
183
184 inferred_model_path = str(temp_path / "onnx_shape_inferred.onnx")
185 onnx.shape_inference.infer_shapes_path(input_model, inferred_model_path)
186 model = onnx.load(inferred_model_path)
187
188 if model is None:
189 model = input_model if isinstance(input_model, onnx.ModelProto) else onnx.load(input_model)
190
191 add_pre_process_metadata(model)
192
193 if save_as_external_data:
194 onnx.save_model(
195 model,
196 output_model_path,
197 save_as_external_data=True,
198 all_tensors_to_one_file=all_tensors_to_one_file,
199 location=external_data_location,
200 size_threshold=external_data_size_threshold,
201 convert_attribute=False,
202 )
203 else:
204 onnx.save(model, output_model_path)
205 