codekingpro/portable-devtools
114k
1from typing import Any, Iterable, Sequence, Type2 3 4from fastembed.common.types import NumpyArray, Device5from fastembed.common import ImageInput, OnnxProvider6from fastembed.common.onnx_model import OnnxOutputContext7from fastembed.common.utils import define_cache_dir, normalize8from fastembed.image.image_embedding_base import ImageEmbeddingBase9from fastembed.image.onnx_image_model import ImageEmbeddingWorker, OnnxImageModel10 11from fastembed.common.model_description import DenseModelDescription, ModelSource12 13supported_onnx_models: list[DenseModelDescription] = [14 DenseModelDescription(15 model="Qdrant/clip-ViT-B-32-vision",16 dim=512,17 description="Image embeddings, Multimodal (text&image), 2021 year",18 license="mit",19 size_in_GB=0.34,20 sources=ModelSource(hf="Qdrant/clip-ViT-B-32-vision"),21 model_file="model.onnx",22 ),23 DenseModelDescription(24 model="Qdrant/resnet50-onnx",25 dim=2048,26 description="Image embeddings, Unimodal (image), 2016 year",27 license="apache-2.0",28 size_in_GB=0.1,29 sources=ModelSource(hf="Qdrant/resnet50-onnx"),30 model_file="model.onnx",31 ),32 DenseModelDescription(33 model="Qdrant/Unicom-ViT-B-16",34 dim=768,35 description="Image embeddings (more detailed than Unicom-ViT-B-32), Multimodal (text&image), 2023 year",36 license="apache-2.0",37 size_in_GB=0.82,38 sources=ModelSource(hf="Qdrant/Unicom-ViT-B-16"),39 model_file="model.onnx",40 ),41 DenseModelDescription(42 model="Qdrant/Unicom-ViT-B-32",43 dim=512,44 description="Image embeddings, Multimodal (text&image), 2023 year",45 license="apache-2.0",46 size_in_GB=0.48,47 sources=ModelSource(hf="Qdrant/Unicom-ViT-B-32"),48 model_file="model.onnx",49 ),50 DenseModelDescription(51 model="jinaai/jina-clip-v1",52 dim=768,53 description="Image embeddings, Multimodal (text&image), 2024 year",54 license="apache-2.0",55 size_in_GB=0.34,56 sources=ModelSource(hf="jinaai/jina-clip-v1"),57 model_file="onnx/vision_model.onnx",58 ),59]60 61 62class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):63 def __init__(64 self,65 model_name: str,66 67 cache_dir: str | None = None,68 threads: int | None = None,69 providers: Sequence[OnnxProvider] | None = None,70 cuda: bool | Device = Device.AUTO,71 device_ids: list[int] | None = None,72 lazy_load: bool = False,73 device_id: int | None = None,74 specific_model_path: str | None = None,75 **kwargs: Any,76 ):77 """78 Args:79 model_name (str): The name of the model to use.80 cache_dir (str, optional): The path to the cache directory.81 Can be set using the `FASTEMBED_CACHE_PATH` env variable.82 Defaults to `fastembed_cache` in the system's temp directory.83 threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.84 providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.85 Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.86 cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`87 Defaults to Device.AUTO.88 device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in89 workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive90 with `providers`. Defaults to None.91 lazy_load (bool, optional): Whether to load the model during class initialization or on demand.92 Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.93 device_id (Optional[int], optional): The device id to use for loading the model in the worker process.94 specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else95 96 Raises:97 ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.98 """99 100 super().__init__(model_name, cache_dir, threads, **kwargs)101 self.providers = providers102 self.lazy_load = lazy_load103 self._extra_session_options = self._select_exposed_session_options(kwargs)104 105 # List of device ids, that can be used for data parallel processing in workers106 self.device_ids = device_ids107 self.cuda = cuda108 109 # This device_id will be used if we need to load model in current process110 self.device_id: int | None = None111 if device_id is not None:112 self.device_id = device_id113 elif self.device_ids is not None:114 self.device_id = self.device_ids[0]115 116 self.model_description = self._get_model_description(model_name)117 self.cache_dir = str(define_cache_dir(cache_dir))118 self._specific_model_path = specific_model_path119 self._model_dir = self.download_model(120 self.model_description,121 self.cache_dir,122 local_files_only=self._local_files_only,123 specific_model_path=self._specific_model_path,124 )125 126 if not self.lazy_load:127 self.load_onnx_model()128 129 def load_onnx_model(self) -> None:130 """131 Load the onnx model.132 """133 self._load_onnx_model(134 model_dir=self._model_dir,135 model_file=self.model_description.model_file,136 threads=self.threads,137 providers=self.providers,138 cuda=self.cuda,139 device_id=self.device_id,140 extra_session_options=self._extra_session_options,141 )142 143 @classmethod144 def _list_supported_models(cls) -> list[DenseModelDescription]:145 """146 Lists the supported models.147 148 Returns:149 list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.150 """151 return supported_onnx_models152 153 def embed(154 self,155 images: ImageInput | Iterable[ImageInput],156 batch_size: int = 16,157 parallel: int | None = None,158 **kwargs: Any,159 ) -> Iterable[NumpyArray]:160 """161 Encode a list of images into list of embeddings.162 We use mean pooling with attention so that the model can handle variable-length inputs.163 164 Args:165 images: Iterator of image paths or single image path to embed166 batch_size: Batch size for encoding -- higher values will use more memory, but be faster167 parallel:168 If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.169 If 0, use all available cores.170 If None, don't use data-parallel processing, use default onnxruntime threading instead.171 172 Returns:173 List of embeddings, one per document174 """175 176 yield from self._embed_images(177 model_name=self.model_name,178 cache_dir=str(self.cache_dir),179 images=images,180 batch_size=batch_size,181 parallel=parallel,182 providers=self.providers,183 cuda=self.cuda,184 device_ids=self.device_ids,185 local_files_only=self._local_files_only,186 specific_model_path=self._specific_model_path,187 extra_session_options=self._extra_session_options,188 **kwargs,189 )190 191 @classmethod192 def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[NumpyArray]"]:193 return OnnxImageEmbeddingWorker194 195 def _preprocess_onnx_input(196 self, onnx_input: dict[str, NumpyArray], **kwargs: Any197 ) -> dict[str, NumpyArray]:198 """199 Preprocess the onnx input.200 """201 202 return onnx_input203 204 def _post_process_onnx_output(205 self, output: OnnxOutputContext, **kwargs: Any206 ) -> Iterable[NumpyArray]:207 return normalize(output.model_output)208 209 210class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):211 def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> OnnxImageEmbedding:212 return OnnxImageEmbedding(213 model_name=model_name,214 cache_dir=cache_dir,215 threads=1,216 **kwargs,217 )218 