codekingpro/portable-devtools
115k
1from typing import Iterable, Any2 3from fastembed.common.model_description import DenseModelDescription4from fastembed.common.types import NumpyArray5from fastembed.common.model_management import ModelManagement6 7 8class TextEmbeddingBase(ModelManagement[DenseModelDescription]):9 def __init__(10 self,11 model_name: str,12 cache_dir: str | None = None,13 threads: int | None = None,14 **kwargs: Any,15 ):16 self.model_name = model_name17 self.cache_dir = cache_dir18 self.threads = threads19 self._local_files_only = kwargs.pop("local_files_only", False)20 self._embedding_size: int | None = None21 22 def embed(23 self,24 documents: str | Iterable[str],25 batch_size: int = 256,26 parallel: int | None = None,27 **kwargs: Any,28 ) -> Iterable[NumpyArray]:29 raise NotImplementedError()30 31 def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:32 """33 Embeds a list of text passages into a list of embeddings.34 35 Args:36 texts (Iterable[str]): The list of texts to embed.37 **kwargs: Additional keyword argument to pass to the embed method.38 39 Yields:40 Iterable[NumpyArray]: The embeddings.41 """42 43 # This is model-specific, so that different models can have specialized implementations44 yield from self.embed(texts, **kwargs)45 46 def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:47 """48 Embeds queries49 50 Args:51 query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.52 53 Returns:54 Iterable[NumpyArray]: The embeddings.55 """56 57 # This is model-specific, so that different models can have specialized implementations58 if isinstance(query, str):59 yield from self.embed([query], **kwargs)60 else:61 yield from self.embed(query, **kwargs)62 63 @classmethod64 def get_embedding_size(cls, model_name: str) -> int:65 """Returns embedding size of the passed model."""66 raise NotImplementedError("Subclasses must implement this method")67 68 @property69 def embedding_size(self) -> int:70 """Returns embedding size for the current model"""71 raise NotImplementedError("Subclasses must implement this method")72 73 def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:74 """Returns the number of tokens in the texts."""75 raise NotImplementedError("Subclasses must implement this method")76 