codekingpro/portable-devtools
115k
1import warnings2from typing import Any, Iterable, Sequence, Type3from dataclasses import asdict4 5from fastembed.common.types import NumpyArray, OnnxProvider, Device6from fastembed.text.clip_embedding import CLIPOnnxEmbedding7from fastembed.text.custom_text_embedding import CustomTextEmbedding8from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding9from fastembed.text.pooled_embedding import PooledEmbedding10from fastembed.text.multitask_embedding import JinaEmbeddingV311from fastembed.text.onnx_embedding import OnnxTextEmbedding12from fastembed.text.text_embedding_base import TextEmbeddingBase13from fastembed.common.model_description import DenseModelDescription, ModelSource, PoolingType14 15 16class TextEmbedding(TextEmbeddingBase):17 EMBEDDINGS_REGISTRY: list[Type[TextEmbeddingBase]] = [18 OnnxTextEmbedding,19 CLIPOnnxEmbedding,20 PooledNormalizedEmbedding,21 PooledEmbedding,22 JinaEmbeddingV3,23 CustomTextEmbedding,24 ]25 26 @classmethod27 def list_supported_models(cls) -> list[dict[str, Any]]:28 """Lists the supported models.29 30 Returns:31 list[dict[str, Any]]: A list of dictionaries containing the model information.32 """33 return [asdict(model) for model in cls._list_supported_models()]34 35 @classmethod36 def _list_supported_models(cls) -> list[DenseModelDescription]:37 result: list[DenseModelDescription] = []38 for embedding in cls.EMBEDDINGS_REGISTRY:39 result.extend(embedding._list_supported_models())40 return result41 42 @classmethod43 def add_custom_model(44 cls,45 model: str,46 pooling: PoolingType,47 normalization: bool,48 sources: ModelSource,49 dim: int,50 model_file: str = "onnx/model.onnx",51 description: str = "",52 license: str = "",53 size_in_gb: float = 0.0,54 additional_files: list[str] | None = None,55 ) -> None:56 registered_models = cls._list_supported_models()57 for registered_model in registered_models:58 if model.lower() == registered_model.model.lower():59 raise ValueError(60 f"Model {model} is already registered in TextEmbedding, if you still want to add this model, "61 f"please use another model name"62 )63 64 CustomTextEmbedding.add_model(65 DenseModelDescription(66 model=model,67 sources=sources,68 dim=dim,69 model_file=model_file,70 description=description,71 license=license,72 size_in_GB=size_in_gb,73 additional_files=additional_files or [],74 ),75 pooling=pooling,76 normalization=normalization,77 )78 79 def __init__(80 self,81 model_name: str = "BAAI/bge-small-en-v1.5",82 cache_dir: str | None = None,83 threads: int | None = None,84 providers: Sequence[OnnxProvider] | None = None,85 cuda: bool | Device = Device.AUTO,86 device_ids: list[int] | None = None,87 lazy_load: bool = False,88 **kwargs: Any,89 ):90 super().__init__(model_name, cache_dir, threads, **kwargs)91 if model_name.lower() == "nomic-ai/nomic-embed-text-v1.5-Q".lower():92 warnings.warn(93 "The model 'nomic-ai/nomic-embed-text-v1.5-Q' has been updated on HuggingFace. Please review "94 "the latest documentation on HF and release notes to ensure compatibility with your workflow. ",95 UserWarning,96 stacklevel=2,97 )98 if model_name.lower() in {99 "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2".lower(),100 "thenlper/gte-large".lower(),101 "intfloat/multilingual-e5-large".lower(),102 "sentence-transformers/paraphrase-multilingual-mpnet-base-v2".lower(),103 }:104 warnings.warn(105 f"The model {model_name} now uses mean pooling instead of CLS embedding. "106 f"In order to preserve the previous behaviour, consider either pinning fastembed version to 0.5.1 or "107 "using `add_custom_model` functionality.",108 UserWarning,109 stacklevel=2,110 )111 for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:112 supported_models = EMBEDDING_MODEL_TYPE._list_supported_models()113 if any(model_name.lower() == model.model.lower() for model in supported_models):114 self.model = EMBEDDING_MODEL_TYPE(115 model_name=model_name,116 cache_dir=cache_dir,117 threads=threads,118 providers=providers,119 cuda=cuda,120 device_ids=device_ids,121 lazy_load=lazy_load,122 **kwargs,123 )124 return125 126 raise ValueError(127 f"Model {model_name} is not supported in TextEmbedding. "128 "Please check the supported models using `TextEmbedding.list_supported_models()`"129 )130 131 @property132 def embedding_size(self) -> int:133 """Get the embedding size of the current model"""134 if self._embedding_size is None:135 self._embedding_size = self.get_embedding_size(self.model_name)136 return self._embedding_size137 138 @classmethod139 def get_embedding_size(cls, model_name: str) -> int:140 """Get the embedding size of the passed model141 142 Args:143 model_name (str): The name of the model to get embedding size for.144 145 Returns:146 int: The size of the embedding.147 148 Raises:149 ValueError: If the model name is not found in the supported models.150 """151 descriptions = cls._list_supported_models()152 embedding_size: int | None = None153 for description in descriptions:154 if description.model.lower() == model_name.lower():155 embedding_size = description.dim156 break157 if embedding_size is None:158 model_names = [description.model for description in descriptions]159 raise ValueError(160 f"Embedding size for model {model_name} was None. "161 f"Available model names: {model_names}"162 )163 return embedding_size164 165 def embed(166 self,167 documents: str | Iterable[str],168 batch_size: int = 256,169 parallel: int | None = None,170 **kwargs: Any,171 ) -> Iterable[NumpyArray]:172 """173 Encode a list of documents into list of embeddings.174 We use mean pooling with attention so that the model can handle variable-length inputs.175 176 Args:177 documents: Iterator of documents or single document to embed178 batch_size: Batch size for encoding -- higher values will use more memory, but be faster179 parallel:180 If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.181 If 0, use all available cores.182 If None, don't use data-parallel processing, use default onnxruntime threading instead.183 184 Returns:185 List of embeddings, one per document186 """187 yield from self.model.embed(documents, batch_size, parallel, **kwargs)188 189 def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:190 """191 Embeds queries192 193 Args:194 query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.195 196 Returns:197 Iterable[NumpyArray]: The embeddings.198 """199 # This is model-specific, so that different models can have specialized implementations200 yield from self.model.query_embed(query, **kwargs)201 202 def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:203 """204 Embeds a list of text passages into a list of embeddings.205 206 Args:207 texts (Iterable[str]): The list of texts to embed.208 **kwargs: Additional keyword argument to pass to the embed method.209 210 Yields:211 Iterable[SparseEmbedding]: The sparse embeddings.212 """213 # This is model-specific, so that different models can have specialized implementations214 yield from self.model.passage_embed(texts, **kwargs)215 216 def token_count(217 self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any218 ) -> int:219 """Returns the number of tokens in the texts.220 221 Args:222 texts (str | Iterable[str]): The list of texts to embed.223 batch_size (int): Batch size for encoding224 225 Returns:226 int: Sum of number of tokens in the texts.227 """228 return self.model.token_count(texts, batch_size=batch_size, **kwargs)229 