Team Ai
Datasetpublic

codekingpro/portable-devtools

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