codekingpro/portable-devtools
114k
1# Copyright (c) Microsoft Corporation. All rights reserved.
2# Licensed under the MIT License.
3
4import argparse
5import copy
6import json
7import sys
8from collections import OrderedDict
9from pprint import pprint
10from typing import Any
11
12import onnx
13
14TuningResults = dict[str, Any]
15
16_TUNING_RESULTS_KEY = "tuning_results"
17
18
19def _find_tuning_results_in_props(metadata_props):
20 for idx, prop in enumerate(metadata_props):
21 if prop.key == _TUNING_RESULTS_KEY:
22 return idx
23 return -1
24
25
26def extract(model: onnx.ModelProto):
27 idx = _find_tuning_results_in_props(model.metadata_props)
28 if idx < 0:
29 return None
30
31 tuning_results_prop = model.metadata_props[idx]
32 return json.loads(tuning_results_prop.value)
33
34
35def embed(model: onnx.ModelProto, tuning_results: list[TuningResults], overwrite=False):
36 idx = _find_tuning_results_in_props(model.metadata_props)
37 assert overwrite or idx <= 0, "the supplied onnx file already have tuning results embedded!"
38
39 if idx >= 0:
40 model.metadata_props.pop(idx)
41
42 entry = model.metadata_props.add()
43 entry.key = _TUNING_RESULTS_KEY
44 entry.value = json.dumps(tuning_results)
45 return model
46
47
48class Merger:
49 class EpAndValidators:
50 def __init__(self, ep: str, validators: dict[str, str]):
51 self.ep = ep
52 self.validators = copy.deepcopy(validators)
53 self.key = (ep, tuple(sorted(validators.items())))
54
55 def __hash__(self):
56 return hash(self.key)
57
58 def __eq__(self, other):
59 return self.ep == other.ep and self.key == other.key
60
61 def __init__(self):
62 self.ev_to_results = OrderedDict()
63
64 def merge(self, tuning_results: list[TuningResults]):
65 for trs in tuning_results:
66 self._merge_one(trs)
67
68 def get_merged(self):
69 tuning_results = []
70 for ev, flat_results in self.ev_to_results.items():
71 results = {}
72 trs = {
73 "ep": ev.ep,
74 "validators": ev.validators,
75 "results": results,
76 }
77 for (op_sig, params_sig), kernel_id in flat_results.items():
78 kernel_map = results.setdefault(op_sig, {})
79 kernel_map[params_sig] = kernel_id
80 tuning_results.append(trs)
81 return tuning_results
82
83 def _merge_one(self, trs: TuningResults):
84 ev = Merger.EpAndValidators(trs["ep"], trs["validators"])
85 flat_results = self.ev_to_results.setdefault(ev, {})
86 for op_sig, kernel_map in trs["results"].items():
87 for params_sig, kernel_id in kernel_map.items():
88 if (op_sig, params_sig) not in flat_results:
89 flat_results[(op_sig, params_sig)] = kernel_id
90
91
92def parse_args():
93 parser = argparse.ArgumentParser()
94 sub_parsers = parser.add_subparsers(help="Command to execute", dest="cmd")
95
96 extract_parser = sub_parsers.add_parser("extract", help="Extract embedded tuning results from an onnx file.")
97 extract_parser.add_argument("input_onnx")
98 extract_parser.add_argument("output_json")
99
100 embed_parser = sub_parsers.add_parser("embed", help="Embed the tuning results into an onnx file.")
101 embed_parser.add_argument("--force", "-f", action="store_true", help="Overwrite the tuning results if it existed.")
102 embed_parser.add_argument("output_onnx", help="Path of the output onnx file.")
103 embed_parser.add_argument("input_onnx", help="Path of the input onnx file.")
104 embed_parser.add_argument("input_json", nargs="+", help="Path(s) of the tuning results file(s) to be embedded.")
105
106 merge_parser = sub_parsers.add_parser("merge", help="Merge multiple tuning results files as a single one.")
107 merge_parser.add_argument("output_json", help="Path of the output tuning results file.")
108 merge_parser.add_argument("input_json", nargs="+", help="Paths of the tuning results files to be merged.")
109
110 pprint_parser = sub_parsers.add_parser("pprint", help="Pretty print the tuning results.")
111 pprint_parser.add_argument("json_or_onnx", help="A tuning results json file or an onnx file.")
112
113 args = parser.parse_args()
114 if len(vars(args)) == 0:
115 parser.print_help()
116 exit(-1)
117 return args
118
119
120def main():
121 args = parse_args()
122 if args.cmd == "extract":
123 tuning_results = extract(onnx.load_model(args.input_onnx))
124 if tuning_results is None:
125 sys.stderr.write(f"{args.input_onnx} does not have tuning results embedded!\n")
126 sys.exit(-1)
127 json.dump(tuning_results, open(args.output_json, "w")) # noqa: SIM115
128 elif args.cmd == "embed":
129 model = onnx.load_model(args.input_onnx)
130 merger = Merger()
131 for tuning_results in [json.load(open(f)) for f in args.input_json]: # noqa: SIM115
132 merger.merge(tuning_results)
133 model = embed(model, merger.get_merged(), args.force)
134 onnx.save_model(model, args.output_onnx)
135 elif args.cmd == "merge":
136 merger = Merger()
137 for tuning_results in [json.load(open(f)) for f in args.input_json]: # noqa: SIM115
138 merger.merge(tuning_results)
139 json.dump(merger.get_merged(), open(args.output_json, "w")) # noqa: SIM115
140 elif args.cmd == "pprint":
141 tuning_results = None
142 try: # noqa: SIM105
143 tuning_results = json.load(open(args.json_or_onnx)) # noqa: SIM115
144 except Exception:
145 # it might be an onnx file otherwise, try it latter
146 pass
147
148 if tuning_results is None:
149 try:
150 model = onnx.load_model(args.json_or_onnx)
151 tuning_results = extract(model)
152 if tuning_results is None:
153 sys.stderr.write(f"{args.input_onnx} does not have tuning results embedded!\n")
154 sys.exit(-1)
155 except Exception:
156 pass
157
158 if tuning_results is None:
159 sys.stderr.write(f"{args.json_or_onnx} is not a valid tuning results file or onnx file!")
160 sys.exit(-1)
161
162 pprint(tuning_results)
163 else:
164 # invalid choice will be handled by the parser
165 pass
166
167
168if __name__ == "__main__":
169 main()
170 