codekingpro/portable-devtools
114k
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 