codekingpro/portable-devtools
114k
1from typing import Any, Iterable, Sequence, Type2 3from fastembed.common.types import NumpyArray, OnnxProvider, Device4from fastembed.common.onnx_model import OnnxOutputContext5from fastembed.common.utils import define_cache_dir, normalize6from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker7from fastembed.text.text_embedding_base import TextEmbeddingBase8from fastembed.common.model_description import DenseModelDescription, ModelSource9 10supported_onnx_models: list[DenseModelDescription] = [11 DenseModelDescription(12 model="BAAI/bge-base-en",13 dim=768,14 description=(15 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "16 "Prefixes for queries/documents: necessary, 2023 year."17 ),18 license="mit",19 size_in_GB=0.42,20 sources=ModelSource(21 hf="Qdrant/fast-bge-base-en",22 url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",23 _deprecated_tar_struct=True,24 ),25 model_file="model_optimized.onnx",26 ),27 DenseModelDescription(28 model="BAAI/bge-base-en-v1.5",29 dim=768,30 description=(31 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "32 "Prefixes for queries/documents: not so necessary, 2023 year."33 ),34 license="mit",35 size_in_GB=0.21,36 sources=ModelSource(37 hf="qdrant/bge-base-en-v1.5-onnx-q",38 url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",39 _deprecated_tar_struct=True,40 ),41 model_file="model_optimized.onnx",42 ),43 DenseModelDescription(44 model="BAAI/bge-large-en-v1.5",45 dim=1024,46 description=(47 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "48 "Prefixes for queries/documents: not so necessary, 2023 year."49 ),50 license="mit",51 size_in_GB=1.20,52 sources=ModelSource(hf="qdrant/bge-large-en-v1.5-onnx"),53 model_file="model.onnx",54 ),55 DenseModelDescription(56 model="BAAI/bge-small-en",57 dim=384,58 description=(59 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "60 "Prefixes for queries/documents: necessary, 2023 year."61 ),62 license="mit",63 size_in_GB=0.13,64 sources=ModelSource(65 hf="Qdrant/bge-small-en",66 url="https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",67 _deprecated_tar_struct=True,68 ),69 model_file="model_optimized.onnx",70 ),71 DenseModelDescription(72 model="BAAI/bge-small-en-v1.5",73 dim=384,74 description=(75 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "76 "Prefixes for queries/documents: not so necessary, 2023 year."77 ),78 license="mit",79 size_in_GB=0.067,80 sources=ModelSource(hf="qdrant/bge-small-en-v1.5-onnx-q"),81 model_file="model_optimized.onnx",82 ),83 DenseModelDescription(84 model="BAAI/bge-small-zh-v1.5",85 dim=512,86 description=(87 "Text embeddings, Unimodal (text), Chinese, 512 input tokens truncation, "88 "Prefixes for queries/documents: not so necessary, 2023 year."89 ),90 license="mit",91 size_in_GB=0.09,92 sources=ModelSource(93 hf="Qdrant/bge-small-zh-v1.5",94 url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",95 _deprecated_tar_struct=True,96 ),97 model_file="model_optimized.onnx",98 ),99 DenseModelDescription(100 model="mixedbread-ai/mxbai-embed-large-v1",101 dim=1024,102 description=(103 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "104 "Prefixes for queries/documents: necessary, 2024 year."105 ),106 license="apache-2.0",107 size_in_GB=0.64,108 sources=ModelSource(hf="mixedbread-ai/mxbai-embed-large-v1"),109 model_file="onnx/model.onnx",110 ),111 DenseModelDescription(112 model="snowflake/snowflake-arctic-embed-xs",113 dim=384,114 description=(115 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "116 "Prefixes for queries/documents: necessary, 2024 year."117 ),118 license="apache-2.0",119 size_in_GB=0.09,120 sources=ModelSource(hf="snowflake/snowflake-arctic-embed-xs"),121 model_file="onnx/model.onnx",122 ),123 DenseModelDescription(124 model="snowflake/snowflake-arctic-embed-s",125 dim=384,126 description=(127 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "128 "Prefixes for queries/documents: necessary, 2024 year."129 ),130 license="apache-2.0",131 size_in_GB=0.13,132 sources=ModelSource(hf="snowflake/snowflake-arctic-embed-s"),133 model_file="onnx/model.onnx",134 ),135 DenseModelDescription(136 model="snowflake/snowflake-arctic-embed-m",137 dim=768,138 description=(139 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "140 "Prefixes for queries/documents: necessary, 2024 year."141 ),142 license="apache-2.0",143 size_in_GB=0.43,144 sources=ModelSource(hf="Snowflake/snowflake-arctic-embed-m"),145 model_file="onnx/model.onnx",146 ),147 DenseModelDescription(148 model="snowflake/snowflake-arctic-embed-m-long",149 dim=768,150 description=(151 "Text embeddings, Unimodal (text), English, 2048 input tokens truncation, "152 "Prefixes for queries/documents: necessary, 2024 year."153 ),154 license="apache-2.0",155 size_in_GB=0.54,156 sources=ModelSource(hf="snowflake/snowflake-arctic-embed-m-long"),157 model_file="onnx/model.onnx",158 ),159 DenseModelDescription(160 model="snowflake/snowflake-arctic-embed-l",161 dim=1024,162 description=(163 "Text embeddings, Unimodal (text), English, 512 input tokens truncation, "164 "Prefixes for queries/documents: necessary, 2024 year."165 ),166 license="apache-2.0",167 size_in_GB=1.02,168 sources=ModelSource(hf="snowflake/snowflake-arctic-embed-l"),169 model_file="onnx/model.onnx",170 ),171 DenseModelDescription(172 model="jinaai/jina-clip-v1",173 dim=768,174 description=(175 "Text embeddings, Multimodal (text&image), English, Prefixes for queries/documents: "176 "not necessary, 2024 year"177 ),178 license="apache-2.0",179 size_in_GB=0.55,180 sources=ModelSource(hf="jinaai/jina-clip-v1"),181 model_file="onnx/text_model.onnx",182 ),183]184 185 186class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):187 """Implementation of the Flag Embedding model."""188 189 @classmethod190 def _list_supported_models(cls) -> list[DenseModelDescription]:191 """192 Lists the supported models.193 194 Returns:195 list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.196 """197 return supported_onnx_models198 199 def __init__(200 self,201 model_name: str = "BAAI/bge-small-en-v1.5",202 cache_dir: str | None = None,203 threads: int | None = None,204 providers: Sequence[OnnxProvider] | None = None,205 cuda: bool | Device = Device.AUTO,206 device_ids: list[int] | None = None,207 lazy_load: bool = False,208 device_id: int | None = None,209 specific_model_path: str | None = None,210 **kwargs: Any,211 ):212 """213 Args:214 model_name (str): The name of the model to use.215 cache_dir (str, optional): The path to the cache directory.216 Can be set using the `FASTEMBED_CACHE_PATH` env variable.217 Defaults to `fastembed_cache` in the system's temp directory.218 threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.219 providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.220 Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.221 cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`222 Defaults to Device.AUTO.223 device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in224 workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive225 with `providers`. Defaults to None.226 lazy_load (bool, optional): Whether to load the model during class initialization or on demand.227 Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.228 device_id (Optional[int], optional): The device id to use for loading the model in the worker process.229 specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else230 231 Raises:232 ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.233 """234 super().__init__(model_name, cache_dir, threads, **kwargs)235 self.providers = providers236 self.lazy_load = lazy_load237 self._extra_session_options = self._select_exposed_session_options(kwargs)238 # List of device ids, that can be used for data parallel processing in workers239 self.device_ids = device_ids240 self.cuda = cuda241 242 # This device_id will be used if we need to load model in current process243 self.device_id: int | None = None244 if device_id is not None:245 self.device_id = device_id246 elif self.device_ids is not None:247 self.device_id = self.device_ids[0]248 249 self.model_description = self._get_model_description(model_name)250 self.cache_dir = str(define_cache_dir(cache_dir))251 self._specific_model_path = specific_model_path252 self._model_dir = self.download_model(253 self.model_description,254 self.cache_dir,255 local_files_only=self._local_files_only,256 specific_model_path=self._specific_model_path,257 )258 259 if not self.lazy_load:260 self.load_onnx_model()261 262 def embed(263 self,264 documents: str | Iterable[str],265 batch_size: int = 256,266 parallel: int | None = None,267 **kwargs: Any,268 ) -> Iterable[NumpyArray]:269 """270 Encode a list of documents into list of embeddings.271 We use mean pooling with attention so that the model can handle variable-length inputs.272 273 Args:274 documents: Iterator of documents or single document to embed275 batch_size: Batch size for encoding -- higher values will use more memory, but be faster276 parallel:277 If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.278 If 0, use all available cores.279 If None, don't use data-parallel processing, use default onnxruntime threading instead.280 281 Returns:282 List of embeddings, one per document283 """284 yield from self._embed_documents(285 model_name=self.model_name,286 cache_dir=str(self.cache_dir),287 documents=documents,288 batch_size=batch_size,289 parallel=parallel,290 providers=self.providers,291 cuda=self.cuda,292 device_ids=self.device_ids,293 local_files_only=self._local_files_only,294 specific_model_path=self._specific_model_path,295 extra_session_options=self._extra_session_options,296 **kwargs,297 )298 299 @classmethod300 def _get_worker_class(cls) -> Type["TextEmbeddingWorker[NumpyArray]"]:301 return OnnxTextEmbeddingWorker302 303 def _preprocess_onnx_input(304 self, onnx_input: dict[str, NumpyArray], **kwargs: Any305 ) -> dict[str, NumpyArray]:306 """307 Preprocess the onnx input.308 """309 return onnx_input310 311 def _post_process_onnx_output(312 self, output: OnnxOutputContext, **kwargs: Any313 ) -> Iterable[NumpyArray]:314 embeddings = output.model_output315 316 if embeddings.ndim == 3: # (batch_size, seq_len, embedding_dim)317 processed_embeddings = embeddings[:, 0]318 elif embeddings.ndim == 2: # (batch_size, embedding_dim)319 processed_embeddings = embeddings320 else:321 raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")322 return normalize(processed_embeddings)323 324 def load_onnx_model(self) -> None:325 self._load_onnx_model(326 model_dir=self._model_dir,327 model_file=self.model_description.model_file,328 threads=self.threads,329 providers=self.providers,330 cuda=self.cuda,331 device_id=self.device_id,332 extra_session_options=self._extra_session_options,333 )334 335 def token_count(336 self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any337 ) -> int:338 return self._token_count(texts, batch_size=batch_size, **kwargs)339 340 341class OnnxTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):342 def init_embedding(343 self,344 model_name: str,345 cache_dir: str,346 **kwargs: Any,347 ) -> OnnxTextEmbedding:348 return OnnxTextEmbedding(349 model_name=model_name,350 cache_dir=cache_dir,351 threads=1,352 **kwargs,353 )354 