()
| 259 | |
| 260 | |
| 261 | def main(): |
| 262 | parser = argparse.ArgumentParser( |
| 263 | description="generate header file for reducing binary size by " |
| 264 | "stripping unused oprs in a particular network; output file would " |
| 265 | "be written to bin_reduce.h", |
| 266 | formatter_class=argparse.ArgumentDefaultsHelpFormatter, |
| 267 | ) |
| 268 | parser.add_argument( |
| 269 | "inputs", |
| 270 | nargs="+", |
| 271 | help="input files that describe specific traits of the network; " |
| 272 | "can be one of the following:" |
| 273 | " 1. json files generated by " |
| 274 | "megbrain.serialize_comp_graph_to_file() in python; " |
| 275 | " 2. trace files generated by midout library", |
| 276 | ) |
| 277 | default_file = os.path.join( |
| 278 | HeaderGen.get_megengine_root(), "src", "bin_reduce_cmake.h" |
| 279 | ) |
| 280 | is_megvii3 = HeaderGen.get_megvii3_root() |
| 281 | if is_megvii3: |
| 282 | default_file = os.path.join( |
| 283 | HeaderGen.get_megvii3_root(), "utils", "bin_reduce.h" |
| 284 | ) |
| 285 | parser.add_argument("-o", "--output", help="output file", default=default_file) |
| 286 | args = parser.parse_args() |
| 287 | print("config output file: {}".format(args.output)) |
| 288 | |
| 289 | gen = HeaderGen() |
| 290 | for i in args.inputs: |
| 291 | print("==== processing {}".format(i)) |
| 292 | with open(i) as fin: |
| 293 | if fin.read(len(MIDOUT_TRACE_MAGIC)) == MIDOUT_TRACE_MAGIC: |
| 294 | gen.extend_midout(i) |
| 295 | gen.extend_elemwise_mode_info(i) |
| 296 | else: |
| 297 | fin.seek(0) |
| 298 | gen.extend_netinfo(json.loads(fin.read())) |
| 299 | |
| 300 | with open(args.output, "w") as fout: |
| 301 | gen.generate(fout) |
| 302 | |
| 303 | |
| 304 | if __name__ == "__main__": |
no test coverage detected