Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
text_embedding.py229 linesDownload Raw Back to text
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 
codekingpro/portable-devtools · Team Ai