Compare commits

...
+101 -4
View File
@@ -3,6 +3,8 @@ import os
from collections import defaultdict
from copy import deepcopy
from enum import Enum
from dataclasses import dataclass
from multiprocessing import shared_memory, Manager, Lock
from multiprocessing import Queue, get_context
from multiprocessing.context import BaseContext
from multiprocessing.process import BaseProcess
@@ -10,6 +12,11 @@ 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
@@ -17,6 +24,13 @@ 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"
@@ -32,12 +46,54 @@ class Worker:
raise NotImplementedError()
class SharedMemoryPool:
def __init__(self, lock: Lock):
self._lock = lock
self._pool: dict[str, tuple[shared_memory.SharedMemory, int, np.dtype]] = {}
self._free_buffers: list[str] = []
def allocate(self, size: int, dtype: np.dtype) -> tuple[shared_memory.SharedMemory, str]:
best_match = None
best_size = float("inf")
for buf_name in self._free_buffers:
shm, buf_size, buf_dtype = self._pool[buf_name]
# get best match for needed size
if buf_size >= size and buf_dtype == dtype and buf_size < best_size:
best_match = buf_name
best_size = buf_size
if best_match:
self._free_buffers.remove(best_match)
return self._pool[best_match][0], best_match
shm = shared_memory.SharedMemory(create=True, size=size)
self._pool[shm.name] = (shm, size, dtype)
return shm, shm.name
# if no match found, create new buffer
shm = shared_memory.SharedMemory(create=True, size=size)
self._pool[shm.name] = (shm, size, dtype)
return shm, shm.name
def release(self, name: str) -> None:
if name in self._pool and name not in self._free_buffers:
self._free_buffers.append(name)
def cleanup(self) -> None:
for shm, _, _ in self._pool.values():
shm.close()
shm.unlink()
self._pool.clear()
self._free_buffers.clear()
def __del__(self):
self.cleanup()
def _worker(
worker_class: Type[Worker],
input_queue: Queue,
output_queue: Queue,
num_active_workers: BaseValue,
worker_id: int,
shared_pool: SharedMemoryPool,
kwargs: Optional[dict[str, Any]] = None,
) -> None:
"""
@@ -55,7 +111,6 @@ def _worker(
try:
worker = worker_class.start(**kwargs)
# Keep going until you get an item that's None.
def input_queue_iterable() -> Iterable[Any]:
while True:
item = input_queue.get()
@@ -64,7 +119,24 @@ def _worker(
yield item
for processed_item in worker.process(input_queue_iterable()):
output_queue.put(processed_item)
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:
shm, shm_name = shared_pool.allocate(array.nbytes, array.dtype)
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))
for field in output_metadata: # mark release to reuse
shared_pool.release(output_metadata[field]["name"])
except Exception as e: # pylint: disable=broad-except
logging.exception(e)
output_queue.put(QueueSignals.error)
@@ -108,6 +180,8 @@ class ParallelWorkerPool:
self.device_ids = device_ids
self.cuda = cuda
self.num_active_workers: Optional[BaseValue] = None
self.manager = Manager()
self.shared_pool = SharedMemoryPool(self.manager.Lock())
def start(self, **kwargs: Any) -> None:
self.input_queue = self.ctx.Queue(self.queue_size)
@@ -133,6 +207,7 @@ class ParallelWorkerPool:
self.output_queue,
self.num_active_workers,
worker_id,
self.shared_pool,
worker_kwargs,
),
)
@@ -178,7 +253,17 @@ class ParallelWorkerPool:
if out_item == QueueSignals.error:
self.join_or_terminate()
raise RuntimeError("Thread unexpectedly terminated")
yield out_item
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()
yield (idx, OnnxOutputContext(**output_arrays))
read += 1
self.input_queue.put((idx, item))
@@ -193,9 +278,18 @@ class ParallelWorkerPool:
if out_item == QueueSignals.error:
self.join_or_terminate()
raise RuntimeError("Thread unexpectedly terminated")
yield out_item
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()
yield (idx, OnnxOutputContext(**output_arrays))
read += 1
finally:
self.shared_pool.cleanup()
assert self.input_queue is not None, "Input queue is None"
assert self.output_queue is not None, "Output queue is None"
self.join()
@@ -231,11 +325,13 @@ class ParallelWorkerPool:
if process.is_alive():
process.terminate()
self.processes.clear()
self.shared_pool.cleanup()
def join(self) -> None:
for process in self.processes:
process.join()
self.processes.clear()
self.shared_pool.cleanup()
def __del__(self) -> None:
"""
@@ -250,3 +346,4 @@ class ParallelWorkerPool:
for process in self.processes:
if process.is_alive():
process.terminate()
self.shared_pool.cleanup()