Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
sparse_text_embedding.py144 linesDownload Raw Back to sparse
1from typing import Any, Iterable, Sequence, Type2from dataclasses import asdict3 4from fastembed.common import OnnxProvider5from fastembed.common.types import Device6from fastembed.sparse.bm25 import Bm257from fastembed.sparse.bm42 import Bm428from fastembed.sparse.minicoil import MiniCOIL9from fastembed.sparse.sparse_embedding_base import (10    SparseEmbedding,11    SparseTextEmbeddingBase,12)13from fastembed.sparse.splade_pp import SpladePP14import warnings15from fastembed.common.model_description import SparseModelDescription16 17 18class SparseTextEmbedding(SparseTextEmbeddingBase):19    EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25, MiniCOIL]20 21    @classmethod22    def list_supported_models(cls) -> list[dict[str, Any]]:23        """24        Lists the supported models.25 26        Returns:27            list[dict[str, Any]]: A list of dictionaries containing the model information.28 29            Example:30                ```31                [32                    {33                        "model": "prithvida/SPLADE_PP_en_v1",34                        "vocab_size": 30522,35                        "description": "Independent Implementation of SPLADE++ Model for English",36                        "license": "apache-2.0",37                        "size_in_GB": 0.532,38                        "sources": {39                            "hf": "qdrant/SPLADE_PP_en_v1",40                        },41                    }42                ]43                ```44        """45        return [asdict(model) for model in cls._list_supported_models()]46 47    @classmethod48    def _list_supported_models(cls) -> list[SparseModelDescription]:49        result: list[SparseModelDescription] = []50        for embedding in cls.EMBEDDINGS_REGISTRY:51            result.extend(embedding._list_supported_models())52        return result53 54    def __init__(55        self,56        model_name: str,57        cache_dir: str | None = None,58        threads: int | None = None,59        providers: Sequence[OnnxProvider] | None = None,60        cuda: bool | Device = Device.AUTO,61        device_ids: list[int] | None = None,62        lazy_load: bool = False,63        **kwargs: Any,64    ):65        super().__init__(model_name, cache_dir, threads, **kwargs)66        if model_name.lower() == "prithvida/Splade_PP_en_v1".lower():67            warnings.warn(68                "The right spelling is prithivida/Splade_PP_en_v1. "69                "Support of this name will be removed soon, please fix the model_name",70                DeprecationWarning,71                stacklevel=2,72            )73            model_name = "prithivida/Splade_PP_en_v1"74 75        for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:76            supported_models = EMBEDDING_MODEL_TYPE._list_supported_models()77            if any(model_name.lower() == model.model.lower() for model in supported_models):78                self.model = EMBEDDING_MODEL_TYPE(79                    model_name,80                    cache_dir,81                    threads=threads,82                    providers=providers,83                    cuda=cuda,84                    device_ids=device_ids,85                    lazy_load=lazy_load,86                    **kwargs,87                )88                return89 90        raise ValueError(91            f"Model {model_name} is not supported in SparseTextEmbedding."92            "Please check the supported models using `SparseTextEmbedding.list_supported_models()`"93        )94 95    def embed(96        self,97        documents: str | Iterable[str],98        batch_size: int = 256,99        parallel: int | None = None,100        **kwargs: Any,101    ) -> Iterable[SparseEmbedding]:102        """103        Encode a list of documents into list of embeddings.104        We use mean pooling with attention so that the model can handle variable-length inputs.105 106        Args:107            documents: Iterator of documents or single document to embed108            batch_size: Batch size for encoding -- higher values will use more memory, but be faster109            parallel:110                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.111                If 0, use all available cores.112                If None, don't use data-parallel processing, use default onnxruntime threading instead.113 114        Returns:115            List of embeddings, one per document116        """117        yield from self.model.embed(documents, batch_size, parallel, **kwargs)118 119    def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:120        """121        Embeds queries122 123        Args:124            query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.125 126        Returns:127            Iterable[SparseEmbedding]: The sparse embeddings.128        """129        yield from self.model.query_embed(query, **kwargs)130 131    def token_count(132        self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any133    ) -> int:134        """Returns the number of tokens in the texts.135 136        Args:137            texts (str | Iterable[str]): The list of texts to embed.138            batch_size (int): Batch size for encoding139 140        Returns:141            int: Sum of number of tokens in the texts.142        """143        return self.model.token_count(texts, batch_size=batch_size, **kwargs)144 
codekingpro/portable-devtools · Team Ai