MCPcopy Create free account
hub / github.com/pytorch/executorch / main

Function main

backends/arm/scripts/evaluate_model.py:289–391  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

287
288
289def 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 )

Callers 1

evaluate_model.pyFile · 0.70

Calls 15

load_calibration_samplesFunction · 0.90
quantize_modelFunction · 0.90
create_partitionerFunction · 0.90
EdgeCompileConfigClass · 0.90
dump_delegation_infoFunction · 0.90
_evaluateFunction · 0.85
infoMethod · 0.80
moduleMethod · 0.80
dumpMethod · 0.80

Tested by

no test coverage detected