| 47 | |
| 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] |
| 76 | |
| 77 | cuts = [] |
| 78 | for i in range(pred.shape[0]): |
| 79 | cuts.extend(self.det.push(pred[i, 25:75, 0], time[i, 25:75])) |
| 80 | return cuts |
| 81 | |
| 82 | def push(self, pixels: np.ndarray, time: np.ndarray): |
| 83 | if self.pixels is None: |
| 84 | self.pixels = pixels |
| 85 | self.time = time |
| 86 | |
| 87 | return self._inference( |
| 88 | np.stack( |
| 89 | ( |
| 90 | np.tile(np.expand_dims(pixels[0], axis=0), (100, 1, 1, 1)), |
| 91 | np.concatenate( |
| 92 | ( |
| 93 | np.tile(np.expand_dims(pixels[0], axis=0), (25, 1, 1, 1)), |
| 94 | pixels[:75], |
| 95 | ), |
| 96 | 0, |
| 97 | ), |
| 98 | ) |
| 99 | ), |
| 100 | np.stack( |
| 101 | ( |
| 102 | np.tile(np.expand_dims(time[0], axis=0), (100,)), |
| 103 | np.concatenate( |
| 104 | (np.tile(np.expand_dims(time[0], axis=0), (25,)), time[:75]), 0 |
| 105 | ), |
| 106 | ) |