(self, output_saved_model_dir, precision="FP32", max_workspace_size_bytes=8000000000, **kwargs)
| 80 | |
| 81 | |
| 82 | def convert(self, output_saved_model_dir, precision="FP32", max_workspace_size_bytes=8000000000, **kwargs): |
| 83 | |
| 84 | if precision == "INT8" and self.calibration_data is None: |
| 85 | raise(Exception("No calibration data set!")) |
| 86 | |
| 87 | trt_precision = precision_dict[precision] |
| 88 | conversion_params = tf_trt.DEFAULT_TRT_CONVERSION_PARAMS._replace(precision_mode=trt_precision, |
| 89 | max_workspace_size_bytes=max_workspace_size_bytes, |
| 90 | use_calibration= precision == "INT8") |
| 91 | converter = tf_trt.TrtGraphConverterV2(input_saved_model_dir=self.input_saved_model_dir, |
| 92 | conversion_params=conversion_params) |
| 93 | |
| 94 | if precision == "INT8": |
| 95 | converter.convert(calibration_input_fn=self.calibration_data) |
| 96 | else: |
| 97 | converter.convert() |
| 98 | |
| 99 | converter.save(output_saved_model_dir=output_saved_model_dir) |
| 100 | |
| 101 | return OptimizedModel(output_saved_model_dir) |
| 102 | |
| 103 | def predict(self, input_data): |
| 104 | if self.loaded_model is None: |
no test coverage detected