codekingpro/portable-devtools
115k
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 