MCPcopy Create free account
hub / github.com/Breakthrough/PySceneDetect / Predictor

Class Predictor

scenedetect/detectors/transnet_v2.py:49–128  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

47
48
49class 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 )

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected