Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_image_model.py157 linesDownload Raw Back to image
1import contextlib2import os3from multiprocessing import get_all_start_methods4from pathlib import Path5from typing import Any, Iterable, Sequence, Type6 7import numpy as np8from PIL import Image9 10from fastembed.image.transform.operators import Compose11from fastembed.common.types import NumpyArray, Device12from fastembed.common import ImageInput, OnnxProvider13from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T14from fastembed.common.preprocessor_utils import load_preprocessor15from fastembed.common.utils import iter_batch16from fastembed.parallel_processor import ParallelWorkerPool17 18# Holds type of the embedding result19 20 21class OnnxImageModel(OnnxModel[T]):22    @classmethod23    def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:24        raise NotImplementedError("Subclasses must implement this method")25 26    def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:27        """Post-process the ONNX model output to convert it into a usable format.28 29        Args:30            output (OnnxOutputContext): The raw output from the ONNX model.31            **kwargs: Additional keyword arguments that may be needed by specific implementations.32 33        Returns:34            Iterable[T]: Post-processed output as an iterable of type T.35        """36        raise NotImplementedError("Subclasses must implement this method")37 38    def __init__(self) -> None:39        super().__init__()40        self.processor: Compose | None = None41 42    def _preprocess_onnx_input(43        self, onnx_input: dict[str, NumpyArray], **kwargs: Any44    ) -> dict[str, NumpyArray]:45        """46        Preprocess the onnx input.47        """48        return onnx_input49 50    def _load_onnx_model(51        self,52        model_dir: Path,53        model_file: str,54        threads: int | None,55        providers: Sequence[OnnxProvider] | None = None,56        cuda: bool | Device = Device.AUTO,57        device_id: int | None = None,58        extra_session_options: dict[str, Any] | None = None,59    ) -> None:60        super()._load_onnx_model(61            model_dir=model_dir,62            model_file=model_file,63            threads=threads,64            providers=providers,65            cuda=cuda,66            device_id=device_id,67            extra_session_options=extra_session_options,68        )69        self.processor = load_preprocessor(model_dir=model_dir)70 71    def load_onnx_model(self) -> None:72        raise NotImplementedError("Subclasses must implement this method")73 74    def _build_onnx_input(self, encoded: NumpyArray) -> dict[str, NumpyArray]:75        input_name = self.model.get_inputs()[0].name  # type: ignore[union-attr]76        return {input_name: encoded}77 78    def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:79        with contextlib.ExitStack() as stack:80            image_files = [81                stack.enter_context(Image.open(image))82                if not isinstance(image, Image.Image)83                else image84                for image in images85            ]86            assert self.processor is not None, "Processor is not initialized"87            encoded = np.array(self.processor(image_files))88        onnx_input = self._build_onnx_input(encoded)89        onnx_input = self._preprocess_onnx_input(onnx_input)90        model_output = self.model.run(None, onnx_input)  # type: ignore[union-attr]91        embeddings = model_output[0].reshape(len(images), -1)92        return OnnxOutputContext(model_output=embeddings)93 94    def _embed_images(95        self,96        model_name: str,97        cache_dir: str,98        images: ImageInput | Iterable[ImageInput],99        batch_size: int = 256,100        parallel: int | None = None,101        providers: Sequence[OnnxProvider] | None = None,102        cuda: bool | Device = Device.AUTO,103        device_ids: list[int] | None = None,104        local_files_only: bool = False,105        specific_model_path: str | None = None,106        extra_session_options: dict[str, Any] | None = None,107        **kwargs: Any,108    ) -> Iterable[T]:109        is_small = False110 111        if isinstance(images, (str, Path, Image.Image)):112            images = [images]113            is_small = True114 115        if isinstance(images, list) and len(images) < batch_size:116            is_small = True117 118        if parallel is None or is_small:119            if not hasattr(self, "model") or self.model is None:120                self.load_onnx_model()121 122            for batch in iter_batch(images, batch_size):123                yield from self._post_process_onnx_output(self.onnx_embed(batch), **kwargs)124        else:125            if parallel == 0:126                parallel = os.cpu_count()127 128            start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"129            params = {130                "model_name": model_name,131                "cache_dir": cache_dir,132                "providers": providers,133                "local_files_only": local_files_only,134                "specific_model_path": specific_model_path,135                **kwargs,136            }137 138            if extra_session_options is not None:139                params.update(extra_session_options)140 141            pool = ParallelWorkerPool(142                num_workers=parallel or 1,143                worker=self._get_worker_class(),144                cuda=cuda,145                device_ids=device_ids,146                start_method=start_method,147            )148            for batch in pool.ordered_map(iter_batch(images, batch_size), **params):149                yield from self._post_process_onnx_output(batch, **kwargs)  # type: ignore150 151 152class ImageEmbeddingWorker(EmbeddingWorker[T]):153    def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:154        for idx, batch in items:155            embeddings = self.model.onnx_embed(batch)156            yield idx, embeddings157 
codekingpro/portable-devtools · Team Ai