diff --git a/fastembed/parallel_processor.py b/fastembed/parallel_processor.py index cb11012..043262d 100644 --- a/fastembed/parallel_processor.py +++ b/fastembed/parallel_processor.py @@ -2,17 +2,41 @@ import logging import os from collections import defaultdict from copy import deepcopy -from multiprocessing import get_context, Condition, Manager +from enum import Enum +from dataclasses import dataclass +from multiprocessing import shared_memory +from multiprocessing import Queue, get_context from multiprocessing.context import BaseContext from multiprocessing.process import BaseProcess -from multiprocessing.sharedctypes import Synchronized -from multiprocessing.synchronize import Lock +from multiprocessing.sharedctypes import Synchronized as BaseValue +from queue import Empty from typing import Any, Iterable, Optional, Type +import numpy as np +from numpy.typing import NDArray + +from fastembed.common.types import NumpyArray + + +# Single item should be processed in less than: +processing_timeout = 10 * 60 # seconds max_internal_batch_size = 200 +@dataclass +class OnnxOutputContext: + model_output: NumpyArray + attention_mask: Optional[NDArray[np.int64]] = None + input_ids: Optional[NDArray[np.int64]] = None + + +class QueueSignals(str, Enum): + stop = "stop" + confirm = "confirm" + error = "error" + + class Worker: @classmethod def start(cls, *args: Any, **kwargs: Any) -> "Worker": @@ -24,57 +48,79 @@ class Worker: def _worker( worker_class: Type[Worker], - task_index: Synchronized, - num_tasks: Synchronized, - input_shared: list[Optional[Any]], - output_shared: list[Optional[Any]], - completed: Synchronized, - lock: Lock, + input_queue: Queue, + output_queue: Queue, + num_active_workers: BaseValue, worker_id: int, kwargs: Optional[dict[str, Any]] = None, ) -> None: + """ + A worker that pulls data pints off the input queue, and places the execution result on the output queue. + When there are no data pints left on the input queue, it decrements + num_active_workers to signal completion. + """ + if kwargs is None: kwargs = {} logging.info( f"Reader worker: {worker_id} PID: {os.getpid()} Device: {kwargs.get('device_id', 'CPU')}" ) + shm_buffers = [] try: - worker: Worker = worker_class.start(**kwargs) - condition = Condition(lock) + worker = worker_class.start(**kwargs) - def task_iterable() -> Iterable[tuple[int, Any]]: + # Keep going until you get an item that's None. + def input_queue_iterable() -> Iterable[Any]: while True: - lock.acquire() - try: - if completed.value >= num_tasks.value and num_tasks.value > 0: - return + item = input_queue.get() + if item == QueueSignals.stop: + break + yield item - if task_index.value >= num_tasks.value: - condition.wait(timeout=0.1) - continue + for processed_item in worker.process(input_queue_iterable()): + idx, output_context = processed_item + output_metadata = {} + for field in ["model_output", "attention_mask", "input_ids"]: + array = getattr(output_context, field, None) + if array is not None: + # Create shared memory and copy array + shm = shared_memory.SharedMemory(create=True, size=array.nbytes) + shm_buffers.append(shm) + shm_array = np.ndarray(array.shape, dtype=array.dtype, buffer=shm.buf) + np.copyto(shm_array, array) + output_metadata[field] = { + "name": shm.name, + "shape": array.shape, + "dtype": array.dtype.str, + } + shm.close() + output_queue.put((idx, output_metadata)) + shm_buffers.clear() - my_task: int = task_index.value - task_index.value += 1 - - batch: Optional[Any] = input_shared[my_task % len(input_shared)] - finally: - lock.release() - - if batch is not None: - yield (my_task, batch) - - for idx, result in worker.process(task_iterable()): - lock.acquire() - try: - output_shared[idx % len(output_shared)] = result - completed.value += 1 - condition.notify_all() - finally: - lock.release() except Exception as e: # pylint: disable=broad-except - logging.exception(f"Reader worker {worker_id} failed: {e}") + logging.exception(e) + output_queue.put(QueueSignals.error) finally: + # It's important that we close and join the queue here before + # decrementing num_active_workers. Otherwise our parent may join us + # before the queue's feeder thread has passed all buffered items to + # the underlying pipe resulting in a deadlock. + # + # See: + # https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#pipes-and-queues + # https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#programming-guidelines + for shm in shm_buffers: + shm.close() + shm.unlink() + input_queue.close() + output_queue.close() + input_queue.join_thread() + output_queue.join_thread() + + with num_active_workers.get_lock(): + num_active_workers.value -= 1 + logging.info(f"Reader worker {worker_id} finished") @@ -89,31 +135,23 @@ class ParallelWorkerPool: ): self.worker_class = worker self.num_workers = num_workers + self.input_queue: Optional[Queue] = None + self.output_queue: Optional[Queue] = None self.ctx: BaseContext = get_context(start_method) self.processes: list[BaseProcess] = [] + self.queue_size = self.num_workers * max_internal_batch_size self.emergency_shutdown = False self.device_ids = device_ids self.cuda = cuda - - self.task_index: Optional[Synchronized[int]] = None - self.num_tasks: Optional[Synchronized[int]] = None - self.completed: Optional[Synchronized[int]] = None - self.lock: Lock - self.condition = None - self.shared_storage_size: int = self.num_workers * max_internal_batch_size - self.input_shared: Optional[list[Optional[Any]]] = None - self.output_shared: Optional[list[Optional[Any]]] = None - self.manager = Manager() + self.num_active_workers: Optional[BaseValue] = None def start(self, **kwargs: Any) -> None: - self.task_index = self.ctx.Value("i", 0) - self.num_tasks = self.ctx.Value("i", 0) - self.completed = self.ctx.Value("i", 0) - self.lock = self.ctx.Lock() - self.condition = self.ctx.Condition(self.lock) + self.input_queue = self.ctx.Queue(self.queue_size) + self.output_queue = self.ctx.Queue(self.queue_size) - self.input_shared = self.manager.list([None] * self.shared_storage_size) - self.output_shared = self.manager.list([None] * self.shared_storage_size) + ctx_value = self.ctx.Value("i", self.num_workers) + assert isinstance(ctx_value, BaseValue) + self.num_active_workers = ctx_value for worker_id in range(0, self.num_workers): worker_kwargs = deepcopy(kwargs) @@ -127,12 +165,9 @@ class ParallelWorkerPool: target=_worker, args=( self.worker_class, - self.task_index, - self.num_tasks, - self.input_shared, - self.output_shared, - self.completed, - self.lock, + self.input_queue, + self.output_queue, + self.num_active_workers, worker_id, worker_kwargs, ), @@ -156,54 +191,78 @@ class ParallelWorkerPool: try: self.start(**kwargs) - total_completed: int = 0 - pushed: int = 0 + assert self.input_queue is not None, "Input queue was not initialized" + assert self.output_queue is not None, "Output queue was not initialized" - for idx, batch in enumerate(stream): + pushed = 0 + read = 0 + for idx, item in enumerate(stream): self.check_worker_health() - self.lock.acquire() - try: - if pushed >= self.shared_storage_size: - while self.completed.value <= total_completed: - self.condition.wait(timeout=0.1) - self.check_worker_health() + if pushed - read < self.queue_size: + try: + out_item = self.output_queue.get_nowait() + except Empty: + out_item = None + else: + try: + out_item = self.output_queue.get(timeout=processing_timeout) + except Empty as e: + self.join_or_terminate() + raise e - self.input_shared[pushed % self.shared_storage_size] = batch - self.num_tasks.value = pushed + 1 - self.condition.notify_all() - pushed += 1 - finally: - self.lock.release() + if out_item is not None: + if out_item == QueueSignals.error: + self.join_or_terminate() + raise RuntimeError("Thread unexpectedly terminated") - self.lock.acquire() - try: - if self.completed.value > total_completed: - for i in range(self.num_tasks.value): - if self.output_shared[i] is not None: - yield (i, self.output_shared[i]) - total_completed += 1 - self.output_shared[i] = None # clear slot (improtant) - finally: - self.lock.release() + idx, output_metadata = out_item + output_arrays = {} + for field, meta in output_metadata.items(): + shm = shared_memory.SharedMemory(name=meta["name"]) + array = np.ndarray( + meta["shape"], dtype=meta["dtype"], buffer=shm.buf + ).copy() + output_arrays[field] = array + shm.close() + shm.unlink() + yield (idx, OnnxOutputContext(**output_arrays)) + read += 1 - while total_completed < pushed: + self.input_queue.put((idx, item)) + pushed += 1 + + for _ in range(self.num_workers): + self.input_queue.put(QueueSignals.stop) + + while read < pushed: self.check_worker_health() - self.lock.acquire() - try: - if self.completed.value > total_completed: - for i in range(self.num_tasks.value): - if self.output_shared[i] is not None: - yield (i, self.output_shared[i]) - total_completed += 1 - self.output_shared[i] = None - elif self.completed.value < pushed: - self.condition.wait(timeout=0.1) - finally: - self.lock.release() - except Exception as e: # pylint: disable=broad-except - logging.exception(f"Error in semi_ordered_map: {e}") + out_item = self.output_queue.get(timeout=processing_timeout) + if out_item == QueueSignals.error: + self.join_or_terminate() + raise RuntimeError("Thread unexpectedly terminated") + + idx, output_metadata = out_item + output_arrays = {} + for field, meta in output_metadata.items(): + shm = shared_memory.SharedMemory(name=meta["name"]) + array = np.ndarray(meta["shape"], dtype=meta["dtype"], buffer=shm.buf).copy() + output_arrays[field] = array + shm.close() + shm.unlink() + yield (idx, OnnxOutputContext(**output_arrays)) + read += 1 finally: + assert self.input_queue is not None, "Input queue is None" + assert self.output_queue is not None, "Output queue is None" self.join() + self.input_queue.close() + self.output_queue.close() + if self.emergency_shutdown: + self.input_queue.cancel_join_thread() + self.output_queue.cancel_join_thread() + else: + self.input_queue.join_thread() + self.output_queue.join_thread() def check_worker_health(self) -> None: """