Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
splade_pp.py197 linesDownload Raw Back to sparse
1from typing import Any, Iterable, Sequence, Type2 3import numpy as np4from fastembed.common import OnnxProvider5from fastembed.common.onnx_model import OnnxOutputContext6from fastembed.common.types import Device7from fastembed.common.utils import define_cache_dir8from fastembed.sparse.sparse_embedding_base import (9    SparseEmbedding,10    SparseTextEmbeddingBase,11)12from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker13from fastembed.common.model_description import SparseModelDescription, ModelSource14 15supported_splade_models: list[SparseModelDescription] = [16    SparseModelDescription(17        model="prithivida/Splade_PP_en_v1",18        vocab_size=30522,19        description="Independent Implementation of SPLADE++ Model for English.",20        license="apache-2.0",21        size_in_GB=0.532,22        sources=ModelSource(hf="Qdrant/Splade_PP_en_v1"),23        model_file="model.onnx",24    ),25    SparseModelDescription(26        model="prithvida/Splade_PP_en_v1",27        vocab_size=30522,28        description="Independent Implementation of SPLADE++ Model for English.",29        license="apache-2.0",30        size_in_GB=0.532,31        sources=ModelSource(hf="Qdrant/Splade_PP_en_v1"),32        model_file="model.onnx",33    ),34]35 36 37class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):38    def _post_process_onnx_output(39        self, output: OnnxOutputContext, **kwargs: Any40    ) -> Iterable[SparseEmbedding]:41        if output.attention_mask is None:42            raise ValueError("attention_mask must be provided for document post-processing")43 44        relu_log = np.log(1 + np.maximum(output.model_output, 0))45 46        weighted_log = relu_log * np.expand_dims(output.attention_mask, axis=-1)47 48        scores = np.max(weighted_log, axis=1)49 50        # Score matrix of shape (batch_size, vocab_size)51        # Most of the values are 0, only a few are non-zero52        for row_scores in scores:53            indices = row_scores.nonzero()[0]54            scores = row_scores[indices]55            yield SparseEmbedding(values=scores, indices=indices)56 57    def token_count(58        self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any59    ) -> int:60        return self._token_count(texts, batch_size=batch_size, **kwargs)61 62    @classmethod63    def _list_supported_models(cls) -> list[SparseModelDescription]:64        """Lists the supported models.65 66        Returns:67            list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.68        """69        return supported_splade_models70 71    def __init__(72        self,73        model_name: str,74        cache_dir: str | None = None,75        threads: int | None = None,76        providers: Sequence[OnnxProvider] | None = None,77        cuda: bool | Device = Device.AUTO,78        device_ids: list[int] | None = None,79        lazy_load: bool = False,80        device_id: int | None = None,81        specific_model_path: str | None = None,82        **kwargs: Any,83    ):84        """85        Args:86            model_name (str): The name of the model to use.87            cache_dir (str, optional): The path to the cache directory.88                                       Can be set using the `FASTEMBED_CACHE_PATH` env variable.89                                       Defaults to `fastembed_cache` in the system's temp directory.90            threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.91            providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.92                Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.93            cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`94                Defaults to Device.95            device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in96                workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive97                with `providers`. Defaults to None.98            lazy_load (bool, optional): Whether to load the model during class initialization or on demand.99                Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.100            device_id (Optional[int], optional): The device id to use for loading the model in the worker process.101            specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else102 103        Raises:104            ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.105        """106        super().__init__(model_name, cache_dir, threads, **kwargs)107        self.providers = providers108        self.lazy_load = lazy_load109        self._extra_session_options = self._select_exposed_session_options(kwargs)110 111        # List of device ids, that can be used for data parallel processing in workers112        self.device_ids = device_ids113        self.cuda = cuda114 115        # This device_id will be used if we need to load model in current process116        self.device_id: int | None = None117        if device_id is not None:118            self.device_id = device_id119        elif self.device_ids is not None:120            self.device_id = self.device_ids[0]121 122        self.model_description = self._get_model_description(model_name)123        self.cache_dir = str(define_cache_dir(cache_dir))124 125        self._specific_model_path = specific_model_path126        self._model_dir = self.download_model(127            self.model_description,128            self.cache_dir,129            local_files_only=self._local_files_only,130            specific_model_path=self._specific_model_path,131        )132 133        if not self.lazy_load:134            self.load_onnx_model()135 136    def load_onnx_model(self) -> None:137        self._load_onnx_model(138            model_dir=self._model_dir,139            model_file=self.model_description.model_file,140            threads=self.threads,141            providers=self.providers,142            cuda=self.cuda,143            device_id=self.device_id,144            extra_session_options=self._extra_session_options,145        )146 147    def embed(148        self,149        documents: str | Iterable[str],150        batch_size: int = 256,151        parallel: int | None = None,152        **kwargs: Any,153    ) -> Iterable[SparseEmbedding]:154        """155        Encode a list of documents into list of embeddings.156        We use mean pooling with attention so that the model can handle variable-length inputs.157 158        Args:159            documents: Iterator of documents or single document to embed160            batch_size: Batch size for encoding -- higher values will use more memory, but be faster161            parallel:162                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.163                If 0, use all available cores.164                If None, don't use data-parallel processing, use default onnxruntime threading instead.165 166        Returns:167            List of embeddings, one per document168        """169        yield from self._embed_documents(170            model_name=self.model_name,171            cache_dir=str(self.cache_dir),172            documents=documents,173            batch_size=batch_size,174            parallel=parallel,175            providers=self.providers,176            cuda=self.cuda,177            device_ids=self.device_ids,178            local_files_only=self._local_files_only,179            specific_model_path=self._specific_model_path,180            extra_session_options=self._extra_session_options,181            **kwargs,182        )183 184    @classmethod185    def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:186        return SpladePPEmbeddingWorker187 188 189class SpladePPEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):190    def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> SpladePP:191        return SpladePP(192            model_name=model_name,193            cache_dir=cache_dir,194            threads=1,195            **kwargs,196        )197 
codekingpro/portable-devtools · Team Ai