From a18b0805138460db5bcac1a29fb0e6cf726cc1d2 Mon Sep 17 00:00:00 2001 From: Yufeng Li Date: Thu, 21 Jul 2022 14:50:28 -0700 Subject: [PATCH] clean up calibration model (#12255) --- .../python/tools/quantization/quantize.py | 22 +++++++++++-------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/onnxruntime/python/tools/quantization/quantize.py b/onnxruntime/python/tools/quantization/quantize.py index 55f1003724..b2d4963fe4 100644 --- a/onnxruntime/python/tools/quantization/quantize.py +++ b/onnxruntime/python/tools/quantization/quantize.py @@ -4,6 +4,7 @@ # license information. # -------------------------------------------------------------------------- import logging +import tempfile from pathlib import Path from onnx import onnx_pb as onnx_proto @@ -142,15 +143,18 @@ def quantize_static( calib_extra_options = { key: extra_options.get(name) for (name, key) in calib_extra_options_keys if name in extra_options } - calibrator = create_calibrator( - model, - op_types_to_quantize, - calibrate_method=calibrate_method, - use_external_data_format=use_external_data_format, - extra_options=calib_extra_options, - ) - calibrator.collect_data(calibration_data_reader) - tensors_range = calibrator.compute_range() + + with tempfile.TemporaryDirectory(prefix="ort.quant.") as quant_tmp_dir: + calibrator = create_calibrator( + model, + op_types_to_quantize, + augmented_model_path=Path(quant_tmp_dir).joinpath("augmented_model.onnx").as_posix(), + calibrate_method=calibrate_method, + use_external_data_format=use_external_data_format, + extra_options=calib_extra_options, + ) + calibrator.collect_data(calibration_data_reader) + tensors_range = calibrator.compute_range() check_static_quant_arguments(quant_format, activation_type, weight_type)