(platform, model_file, weight_file, input_file, mace_out_file,
input_shape, output_shape, input_data_format,
output_data_format, input_node, output_node,
validation_threshold, input_data_type, backend,
validation_outputs_data, log_file)
| 519 | |
| 520 | |
| 521 | def validate(platform, model_file, weight_file, input_file, mace_out_file, |
| 522 | input_shape, output_shape, input_data_format, |
| 523 | output_data_format, input_node, output_node, |
| 524 | validation_threshold, input_data_type, backend, |
| 525 | validation_outputs_data, log_file): |
| 526 | if not isinstance(validation_outputs_data, list): |
| 527 | if os.path.isfile(validation_outputs_data): |
| 528 | validation_outputs = [validation_outputs_data] |
| 529 | else: |
| 530 | validation_outputs = [] |
| 531 | else: |
| 532 | validation_outputs = validation_outputs_data |
| 533 | if validation_outputs: |
| 534 | validate_with_file(platform, output_node, output_shape, |
| 535 | mace_out_file, validation_outputs, |
| 536 | validation_threshold, log_file, |
| 537 | output_data_format) |
| 538 | elif platform == Platform.TENSORFLOW: |
| 539 | validate_tf_model(platform, model_file, input_file, mace_out_file, |
| 540 | input_node, input_shape, input_data_format, |
| 541 | output_node, output_shape, output_data_format, |
| 542 | validation_threshold, input_data_type, |
| 543 | log_file) |
| 544 | elif platform == Platform.PYTORCH: |
| 545 | validate_pytorch_model(platform, model_file, input_file, mace_out_file, |
| 546 | input_node, input_shape, input_data_format, |
| 547 | output_node, output_shape, output_data_format, |
| 548 | validation_threshold, input_data_type, |
| 549 | log_file) |
| 550 | elif platform == Platform.CAFFE: |
| 551 | validate_caffe_model(platform, model_file, |
| 552 | input_file, mace_out_file, weight_file, |
| 553 | input_node, input_shape, input_data_format, |
| 554 | output_node, output_shape, output_data_format, |
| 555 | validation_threshold, log_file) |
| 556 | elif platform == Platform.ONNX: |
| 557 | validate_onnx_model(platform, model_file, |
| 558 | input_file, mace_out_file, |
| 559 | input_node, input_shape, input_data_format, |
| 560 | output_node, output_shape, output_data_format, |
| 561 | validation_threshold, |
| 562 | input_data_type, backend, log_file) |
| 563 | elif platform == Platform.MEGENGINE: |
| 564 | validate_megengine_model(platform, model_file, |
| 565 | input_file, mace_out_file, |
| 566 | input_node, input_shape, |
| 567 | input_data_format, |
| 568 | output_node, output_shape, |
| 569 | output_data_format, |
| 570 | validation_threshold, |
| 571 | input_data_type, log_file) |
| 572 | elif platform == Platform.KERAS: |
| 573 | validate_keras_model(platform, model_file, input_file, mace_out_file, |
| 574 | input_node, input_shape, input_data_format, |
| 575 | output_node, output_shape, output_data_format, |
| 576 | validation_threshold, input_data_type, |
| 577 | log_file) |
| 578 | else: |
nothing calls this directly
no test coverage detected