()
| 287 | |
| 288 | |
| 289 | def main() -> None: |
| 290 | try: |
| 291 | args = _get_args() |
| 292 | except ValueError as e: |
| 293 | logging.error(f"Argument error: {e}") |
| 294 | sys.exit(1) |
| 295 | |
| 296 | # if we have custom ops, register them before processing the model |
| 297 | if args.so_library is not None: |
| 298 | logging.info(f"Loading custom ops from {args.so_library}") |
| 299 | torch.ops.load_library(args.so_library) |
| 300 | |
| 301 | # Get the model and its example inputs |
| 302 | original_model, example_inputs = get_model_and_inputs_from_name( |
| 303 | args.model_name, None |
| 304 | ) |
| 305 | |
| 306 | # Use original model as reference to compare against |
| 307 | ref_model = original_model.eval() |
| 308 | eval_model = ref_model |
| 309 | eval_inputs = example_inputs |
| 310 | |
| 311 | # Cast model and inputs to eval_dtype if specified |
| 312 | if args.dtype is not None: |
| 313 | eval_dtype = _DTYPE_MAP[args.dtype] |
| 314 | eval_model = copy.deepcopy(original_model).to(eval_dtype).eval() |
| 315 | eval_inputs = tuple( |
| 316 | inp.to(eval_dtype) if isinstance(inp, torch.Tensor) else inp |
| 317 | for inp in example_inputs |
| 318 | ) |
| 319 | |
| 320 | # Export the model |
| 321 | exported_program = torch.export.export(eval_model, eval_inputs) |
| 322 | |
| 323 | model_name = os.path.basename(os.path.splitext(args.model_name)[0]) |
| 324 | if args.intermediates: |
| 325 | os.makedirs(args.intermediates, exist_ok=True) |
| 326 | |
| 327 | # We only support Python3.10 and above, so use a later pickle protocol |
| 328 | torch.export.save( |
| 329 | exported_program, |
| 330 | f"{args.intermediates}/{model_name}_exported_program.pt2", |
| 331 | pickle_protocol=5, |
| 332 | ) |
| 333 | |
| 334 | compile_spec = _get_compile_spec(args) |
| 335 | |
| 336 | # Quantize the model if requested |
| 337 | if args.quant_mode is not None: |
| 338 | calibration_samples = None |
| 339 | if ( |
| 340 | "imagenet" in args.evaluators |
| 341 | and args.calibration_data is not None |
| 342 | and Path(args.calibration_data).is_dir() |
| 343 | ): |
| 344 | calibration_samples = _build_imagenet_calibration_samples( |
| 345 | args.calibration_data, CALIBRATION_MAX_SAMPLES |
| 346 | ) |
no test coverage detected