Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
parallel_processor.py254 linesDownload Raw Back to fastembed
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 
codekingpro/portable-devtools · Team Ai