Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
text_embedding_base.py76 linesDownload Raw Back to text
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 
codekingpro/portable-devtools · Team Ai