codekingpro/portable-devtools
115k
1from dataclasses import dataclass2from typing import Iterable, Any3 4import numpy as np5from numpy.typing import NDArray6 7from fastembed.common.model_description import SparseModelDescription8from fastembed.common.types import NumpyArray9from fastembed.common.model_management import ModelManagement10 11 12@dataclass13class SparseEmbedding:14 values: NumpyArray15 indices: NDArray[np.int64] | NDArray[np.int32]16 17 def as_object(self) -> dict[str, NumpyArray]:18 return {19 "values": self.values,20 "indices": self.indices,21 }22 23 def as_dict(self) -> dict[int, float]:24 return {int(i): float(v) for i, v in zip(self.indices, self.values)} # type: ignore25 26 @classmethod27 def from_dict(cls, data: dict[int, float]) -> "SparseEmbedding":28 if len(data) == 0:29 return cls(values=np.array([]), indices=np.array([]))30 indices, values = zip(*data.items())31 return cls(values=np.array(values), indices=np.array(indices))32 33 34class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):35 def __init__(36 self,37 model_name: str,38 cache_dir: str | None = None,39 threads: int | None = None,40 **kwargs: Any,41 ):42 self.model_name = model_name43 self.cache_dir = cache_dir44 self.threads = threads45 self._local_files_only = kwargs.pop("local_files_only", False)46 47 def embed(48 self,49 documents: str | Iterable[str],50 batch_size: int = 256,51 parallel: int | None = None,52 **kwargs: Any,53 ) -> Iterable[SparseEmbedding]:54 raise NotImplementedError()55 56 def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:57 """58 Embeds a list of text passages into a list of embeddings.59 60 Args:61 texts (Iterable[str]): The list of texts to embed.62 **kwargs: Additional keyword argument to pass to the embed method.63 64 Yields:65 Iterable[SparseEmbedding]: The sparse embeddings.66 """67 68 # This is model-specific, so that different models can have specialized implementations69 yield from self.embed(texts, **kwargs)70 71 def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:72 """73 Embeds queries74 75 Args:76 query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.77 78 Returns:79 Iterable[SparseEmbedding]: The sparse embeddings.80 """81 82 # This is model-specific, so that different models can have specialized implementations83 if isinstance(query, str):84 yield from self.embed([query], **kwargs)85 else:86 yield from self.embed(query, **kwargs)87 88 def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:89 """Returns the number of tokens in the texts."""90 raise NotImplementedError("Subclasses must implement this method")91 