| 48 | |
| 49 | class Predictor: |
| 50 | def __init__( |
| 51 | self, |
| 52 | model_path: str | Path, |
| 53 | flash_filter: FlashFilter, |
| 54 | onnx_providers: list[str] | None, |
| 55 | threshold, |
| 56 | ): |
| 57 | import onnxruntime as ort # pyright: ignore[reportMissingImports] |
| 58 | |
| 59 | ort.set_default_logger_severity(3) |
| 60 | |
| 61 | if onnx_providers is None: |
| 62 | onnx_providers = ort.get_available_providers() |
| 63 | |
| 64 | sess_opt = ort.SessionOptions() |
| 65 | sess_opt.log_severity_level = 3 |
| 66 | |
| 67 | self.session = ort.InferenceSession(model_path, sess_opt=sess_opt, providers=onnx_providers) |
| 68 | |
| 69 | self.pixels = None |
| 70 | self.time = None |
| 71 | |
| 72 | self.det = Detector(threshold, flash_filter) |
| 73 | |
| 74 | def _inference(self, pixels: np.ndarray, time: np.ndarray): |
| 75 | pred = np.array(self.session.run(["output"], {"input": pixels}))[0] |