codekingpro/portable-devtools
114k
1# Copyright (c) Microsoft Corporation. All rights reserved.
2# Licensed under the MIT License.
3
4import pathlib
5import typing
6
7from ..logger import get_logger
8from .operator_type_usage_processors import OperatorTypeUsageManager
9from .ort_model_processor import OrtFormatModelProcessor
10
11log = get_logger("ort_format_model.utils")
12
13
14def _extract_ops_and_types_from_ort_models(model_files: typing.Iterable[pathlib.Path], enable_type_reduction: bool):
15 required_ops = {}
16 op_type_usage_manager = OperatorTypeUsageManager() if enable_type_reduction else None
17
18 for model_file in model_files:
19 if not model_file.is_file():
20 raise ValueError(f"Path is not a file: '{model_file}'")
21 model_processor = OrtFormatModelProcessor(str(model_file), required_ops, op_type_usage_manager)
22 model_processor.process() # this updates required_ops and op_type_processors
23
24 return required_ops, op_type_usage_manager
25
26
27def create_config_from_models(
28 model_files: typing.Iterable[pathlib.Path], output_file: pathlib.Path, enable_type_reduction: bool
29):
30 """
31 Create a configuration file with required operators and optionally required types.
32 :param model_files: Model files to use to generate the configuration file.
33 :param output_file: File to write configuration to.
34 :param enable_type_reduction: Include required type information for individual operators in the configuration.
35 """
36
37 required_ops, op_type_processors = _extract_ops_and_types_from_ort_models(model_files, enable_type_reduction)
38
39 output_file.parent.mkdir(parents=True, exist_ok=True)
40
41 with open(output_file, "w") as out:
42 out.write("# Generated from model/s:\n")
43 out.writelines(f"# - {model_file}\n" for model_file in sorted(model_files))
44
45 for domain in sorted(required_ops.keys()):
46 for opset in sorted(required_ops[domain].keys()):
47 ops = required_ops[domain][opset]
48 if ops:
49 out.write(f"{domain};{opset};")
50 if enable_type_reduction:
51 # type string is empty if op hasn't been seen
52 entries = [
53 "{}{}".format(op, op_type_processors.get_config_entry(domain, op) or "")
54 for op in sorted(ops)
55 ]
56 else:
57 entries = sorted(ops)
58
59 out.write("{}\n".format(",".join(entries)))
60
61 log.info("Created config in %s", output_file)
62 