Team Ai
Datasetpublic

codekingpro/portable-devtools

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