mirror of
https://github.com/qdrant/fastembed.git
synced 2026-09-22 22:17:49 -05:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d74e5c633f | ||
|
|
173d82cea5 | ||
|
|
a9ecce9b8a | ||
|
|
673fd6fc0e | ||
|
|
8b789045bd |
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user