diff --git a/machine-learning/immich_ml/sessions/rknn/__init__.py b/machine-learning/immich_ml/sessions/rknn/__init__.py index 721fe5bbda9c3d..0d95e6cf9463ae 100644 --- a/machine-learning/immich_ml/sessions/rknn/__init__.py +++ b/machine-learning/immich_ml/sessions/rknn/__init__.py @@ -64,8 +64,7 @@ def run( run_options: Any = None, ) -> list[NDArray[np.float32]]: input_data: list[NDArray[np.float32]] = [np.ascontiguousarray(v) for v in input_feed.values()] - self.rknnpool.put(input_data) - res = self.rknnpool.get() + res = self.rknnpool.run(input_data) if res is None: raise RuntimeError("RKNN inference failed!") return res diff --git a/machine-learning/immich_ml/sessions/rknn/rknnpool.py b/machine-learning/immich_ml/sessions/rknn/rknnpool.py index fd0af8bcc4e0a7..1caec1abba4f37 100644 --- a/machine-learning/immich_ml/sessions/rknn/rknnpool.py +++ b/machine-learning/immich_ml/sessions/rknn/rknnpool.py @@ -2,9 +2,9 @@ # Following Apache License 2.0 import logging -from concurrent.futures import Future, ThreadPoolExecutor +import threading +from concurrent.futures import ThreadPoolExecutor from pathlib import Path -from queue import Queue from typing import Callable import numpy as np @@ -66,20 +66,18 @@ def __init__( func: Callable[["RKNNLite", list[NDArray[np.float32]]], list[NDArray[np.float32]]], ) -> None: self.tpes = tpes - self.queue: Queue[Future[list[NDArray[np.float32]]]] = Queue() self.rknn_pool = [init_rknn(model_path) for _ in range(tpes)] self.pool = ThreadPoolExecutor(max_workers=tpes) self.func = func self.num = 0 + self.lock = threading.Lock() - def put(self, inputs: list[NDArray[np.float32]]) -> None: - self.queue.put(self.pool.submit(self.func, self.rknn_pool[self.num % self.tpes], inputs)) - self.num += 1 + def run(self, inputs: list[NDArray[np.float32]]) -> list[NDArray[np.float32]]: + with self.lock: + idx = self.num % self.tpes + self.num += 1 - def get(self) -> list[NDArray[np.float32]] | None: - if self.queue.empty(): - return None - fut = self.queue.get() + fut = self.pool.submit(self.func, self.rknn_pool[idx], inputs) return fut.result() def release(self) -> None: diff --git a/machine-learning/test_main.py b/machine-learning/test_main.py index 9e3fe50cba5f7d..a50bea00b734c1 100644 --- a/machine-learning/test_main.py +++ b/machine-learning/test_main.py @@ -564,7 +564,7 @@ def test_run_rknn(self, rknn_session: mock.Mock, mocker: MockerFixture) -> None: session.run(None, input_feed) - rknn_session.return_value.put.assert_called_once_with([input1, input2]) + rknn_session.return_value.run.assert_called_once_with([input1, input2]) assert np_spy.call_count == 2 np_spy.assert_has_calls([mock.call(input1), mock.call(input2)])