2from .
import ConfigABC, QuantConfig, is_valid_dir_path, check_save_path
16 "calib_data_path":
"calibData",
17 "save_path":
"savePath",
18 "model_nickname":
"modelNickname",
19 "save_msgpack_name":
"saveMsgpackName",
20 "model_path":
"modelPath",
21 "inference_scheme":
"inferenceScheme",
22 "save_sample":
"saveSample",
23 "cpu_offload":
"cpuOffload",
24 "singlecore_compile":
"singlecoreCompile",
25 "optimize_option":
"optimizeOption",
26 "sample_dtype":
"sampleDtype",
27 "preprocess_dict":
"preprocess",
30 min_max_check = {
"optimize_option": [0,
None]}
31 inference_scheme_types = [
"single",
"multi",
"global",
"global4",
"global8"]
32 sample_dtypes = [
"float",
"int8"]
37 calib_data_path: str |
None,
38 use_random_calib: bool,
41 save_msgpack_name: str |
None,
42 model_path: str |
None,
43 inference_scheme: str,
46 singlecore_compile: bool,
49 preprocess_dict: dict |
None,
52 quant_config: QuantConfig,
53 optional: bool =
False,
55 super().__init__(optional)
57 calib_data_path, use_random_calib
74 except_list=[
"optional",
"quant_config",
"calib_file",
"buffer_mode"]
78 def default_config(cls):
79 calib_data_path =
None
80 use_random_calib =
False
81 save_path =
"./tmp.mxq"
82 save_msgpack_name =
None
83 model_nickname =
"temporary"
85 inference_scheme =
"single"
88 singlecore_compile =
False
90 sample_dtype =
"float"
91 preprocess_dict =
None
93 quant_config = QuantConfig.default_config()
98 calib_data_path=calib_data_path,
99 use_random_calib=use_random_calib,
101 save_msgpack_name=save_msgpack_name,
102 model_nickname=model_nickname,
103 model_path=model_path,
104 inference_scheme=inference_scheme,
105 save_sample=save_sample,
106 cpu_offload=cpu_offload,
107 singlecore_compile=singlecore_compile,
108 optimize_option=optimize_option,
109 sample_dtype=sample_dtype,
110 preprocess_dict=preprocess_dict,
111 buffer_mode=buffer_mode,
113 quant_config=quant_config,
116 def check_valid_str_configs(self):
121 if self.
inference_scheme.lower()
not in CompileConfig.inference_scheme_types:
123 f
"Inference_scheme should be one of the {CompileConfig.inference_scheme_types}"
125 if self.
sample_dtype.lower()
not in CompileConfig.sample_dtypes:
127 f
"Sample data type should be one of the {CompileConfig.sample_dtypes}"
130 raise ValueError(f
"buffer mode must be either 0 or 1")
133 compile_config_dict = dict()
134 for py_name, cpp_name
in CompileConfig.py_to_cpp_name.items():
135 value = getattr(self, py_name)
136 if isinstance(value, str)
and py_name
in [
140 value = value.lower()
141 compile_config_dict[cpp_name] = value
142 compile_config_dict[
"quant"] = self.
quant_config.to_dict()
144 return compile_config_dict
146 def _get_calibration_path(self, calib_data_path: str, use_random_calib: bool):
147 if not calib_data_path
and not use_random_calib:
149 "Please use calib_data_path or enable use_random_calib. Do not leave calib_data_path empty while setting use_random_calib to false."
152 not calib_data_path
or use_random_calib
155 calib_data_case = check_calib_data(calib_data_path)
156 if calib_data_case == CalibType.SINGLE_DIR:
157 self.
calib_file = tempfile.NamedTemporaryFile(suffix=
".txt")
158 list_np_files_in_txt(calib_data_path, self.
calib_file.name)
160 elif calib_data_case == CalibType.MULTI_DIR:
161 self.
calib_file = tempfile.NamedTemporaryFile(suffix=
".json")
162 list_np_files_in_json(calib_data_path, self.
calib_file.name)
164 elif calib_data_case
in (CalibType.SINGLE_TXT, CalibType.MULTI_JSON):
167 raise ValueError(f
"Got unexpected calib_data_path={calib_data_path}.")
169 return calib_data_path