codekingpro/portable-devtools
115k
1import logging2import os3from collections import defaultdict4from copy import deepcopy5from enum import Enum6from multiprocessing import Queue, get_context7from multiprocessing.context import BaseContext8from multiprocessing.process import BaseProcess9from multiprocessing.sharedctypes import Synchronized as BaseValue10from queue import Empty11from typing import Any, Iterable, Type12 13from fastembed.common.types import Device14 15# Single item should be processed in less than:16processing_timeout = 10 * 60 # seconds17 18max_internal_batch_size = 20019 20 21class QueueSignals(str, Enum):22 stop = "stop"23 confirm = "confirm"24 error = "error"25 26 27class Worker:28 @classmethod29 def start(cls, *args: Any, **kwargs: Any) -> "Worker":30 raise NotImplementedError()31 32 def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:33 raise NotImplementedError()34 35 36def _worker(37 worker_class: Type[Worker],38 input_queue: Queue,39 output_queue: Queue,40 num_active_workers: BaseValue,41 worker_id: int,42 kwargs: dict[str, Any] | None = None,43) -> None:44 """45 A worker that pulls data pints off the input queue, and places the execution result on the output queue.46 When there are no data pints left on the input queue, it decrements47 num_active_workers to signal completion.48 """49 50 if kwargs is None:51 kwargs = {}52 53 logging.info(54 f"Reader worker: {worker_id} PID: {os.getpid()} Device: {kwargs.get('device_id', 'CPU')}"55 )56 try:57 worker = worker_class.start(**kwargs)58 59 # Keep going until you get an item that's None.60 def input_queue_iterable() -> Iterable[Any]:61 while True:62 item = input_queue.get()63 if item == QueueSignals.stop:64 break65 yield item66 67 for processed_item in worker.process(input_queue_iterable()):68 output_queue.put(processed_item)69 except Exception as e: # pylint: disable=broad-except70 logging.exception(e)71 output_queue.put(QueueSignals.error)72 finally:73 # It's important that we close and join the queue here before74 # decrementing num_active_workers. Otherwise our parent may join us75 # before the queue's feeder thread has passed all buffered items to76 # the underlying pipe resulting in a deadlock.77 #78 # See:79 # https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#pipes-and-queues80 # https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#programming-guidelines81 input_queue.close()82 output_queue.close()83 input_queue.join_thread()84 output_queue.join_thread()85 86 with num_active_workers.get_lock():87 num_active_workers.value -= 188 89 logging.info(f"Reader worker {worker_id} finished")90 91 92class ParallelWorkerPool:93 def __init__(94 self,95 num_workers: int,96 worker: Type[Worker],97 start_method: str | None = None,98 device_ids: list[int] | None = None,99 cuda: bool | Device = Device.AUTO,100 ):101 self.worker_class = worker102 self.num_workers = num_workers103 self.input_queue: Queue | None = None104 self.output_queue: Queue | None = None105 self.ctx: BaseContext = get_context(start_method)106 self.processes: list[BaseProcess] = []107 self.queue_size = self.num_workers * max_internal_batch_size108 self.emergency_shutdown = False109 self.device_ids = device_ids110 self.cuda = cuda111 self.num_active_workers: BaseValue | None = None112 113 def start(self, **kwargs: Any) -> None:114 self.input_queue = self.ctx.Queue(self.queue_size)115 self.output_queue = self.ctx.Queue(self.queue_size)116 117 ctx_value = self.ctx.Value("i", self.num_workers)118 assert isinstance(ctx_value, BaseValue)119 self.num_active_workers = ctx_value120 121 for worker_id in range(0, self.num_workers):122 worker_kwargs = deepcopy(kwargs)123 if self.device_ids:124 device_id = self.device_ids[worker_id % len(self.device_ids)]125 worker_kwargs["device_id"] = device_id126 worker_kwargs["cuda"] = self.cuda127 128 assert hasattr(self.ctx, "Process")129 process = self.ctx.Process(130 target=_worker,131 args=(132 self.worker_class,133 self.input_queue,134 self.output_queue,135 self.num_active_workers,136 worker_id,137 worker_kwargs,138 ),139 )140 process.start()141 self.processes.append(process)142 143 def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Any]:144 buffer: defaultdict[int, Any] = defaultdict(Any) # type: ignore145 next_expected = 0146 147 for idx, item in self.semi_ordered_map(stream, *args, **kwargs):148 buffer[idx] = item149 while next_expected in buffer:150 yield buffer.pop(next_expected)151 next_expected += 1152 153 def semi_ordered_map(154 self, stream: Iterable[Any], *args: Any, **kwargs: Any155 ) -> Iterable[tuple[int, Any]]:156 try:157 self.start(**kwargs)158 159 assert self.input_queue is not None, "Input queue was not initialized"160 assert self.output_queue is not None, "Output queue was not initialized"161 162 pushed = 0163 read = 0164 for idx, item in enumerate(stream):165 self.check_worker_health()166 if pushed - read < self.queue_size:167 try:168 out_item = self.output_queue.get_nowait()169 except Empty:170 out_item = None171 else:172 try:173 out_item = self.output_queue.get(timeout=processing_timeout)174 except Empty as e:175 self.join_or_terminate()176 raise e177 178 if out_item is not None:179 if out_item == QueueSignals.error:180 self.join_or_terminate()181 raise RuntimeError("Thread unexpectedly terminated")182 yield out_item183 read += 1184 185 self.input_queue.put((idx, item))186 pushed += 1187 188 for _ in range(self.num_workers):189 self.input_queue.put(QueueSignals.stop)190 191 while read < pushed:192 self.check_worker_health()193 out_item = self.output_queue.get(timeout=processing_timeout)194 if out_item == QueueSignals.error:195 self.join_or_terminate()196 raise RuntimeError("Thread unexpectedly terminated")197 yield out_item198 read += 1199 finally:200 assert self.input_queue is not None, "Input queue is None"201 assert self.output_queue is not None, "Output queue is None"202 self.join()203 self.input_queue.close()204 self.output_queue.close()205 if self.emergency_shutdown:206 self.input_queue.cancel_join_thread()207 self.output_queue.cancel_join_thread()208 else:209 self.input_queue.join_thread()210 self.output_queue.join_thread()211 212 def check_worker_health(self) -> None:213 """214 Checks if any worker process has terminated unexpectedly215 """216 for process in self.processes:217 if not process.is_alive() and process.exitcode != 0:218 self.emergency_shutdown = True219 self.join_or_terminate()220 raise RuntimeError(221 f"Worker PID: {process.pid} terminated unexpectedly with code {process.exitcode}"222 )223 224 def join_or_terminate(self, timeout: int = 1) -> None:225 """226 Emergency shutdown227 @param timeout:228 @return:229 """230 for process in self.processes:231 process.join(timeout=timeout)232 if process.is_alive():233 process.terminate()234 self.processes.clear()235 236 def join(self) -> None:237 for process in self.processes:238 process.join()239 self.processes.clear()240 241 def __del__(self) -> None:242 """243 Terminate processes if the user hasn't joined. This is necessary as244 leaving stray processes running can corrupt shared state. In brief,245 we've observed shared memory counters being reused (when the memory was246 free from the perspective of the parent process) while the stray247 workers still held a reference to them.248 For a discussion of using destructors in Python in this manner, see249 https://eli.thegreenplace.net/2009/06/12/safely-using-destructors-in-python/.250 """251 for process in self.processes:252 if process.is_alive():253 process.terminate()254 