Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
infinity_local.py158 linesDownload Raw Back to embeddings
1"""written under MIT Licence, Michael Feil 2023."""2 3import asyncio4from logging import getLogger5from typing import Any, List, Optional6 7from langchain_core.embeddings import Embeddings8from pydantic import BaseModel, ConfigDict, model_validator9from typing_extensions import Self10 11__all__ = ["InfinityEmbeddingsLocal"]12 13logger = getLogger(__name__)14 15 16class InfinityEmbeddingsLocal(BaseModel, Embeddings):17    """Optimized Infinity embedding models.18 19    https://github.com/michaelfeil/infinity20    This class deploys a local Infinity instance to embed text.21    The class requires async usage.22 23    Infinity is a class to interact with Embedding Models on https://github.com/michaelfeil/infinity24 25 26    Example:27        .. code-block:: python28 29            from langchain_community.embeddings import InfinityEmbeddingsLocal30            async with InfinityEmbeddingsLocal(31                model="BAAI/bge-small-en-v1.5",32                revision=None,33                device="cpu",34            ) as embedder:35                embeddings = await engine.aembed_documents(["text1", "text2"])36    """37 38    model: str39    "Underlying model id from huggingface, e.g. BAAI/bge-small-en-v1.5"40 41    revision: Optional[str] = None42    "Model version, the commit hash from huggingface"43 44    batch_size: int = 3245    "Internal batch size for inference, e.g. 32"46 47    device: str = "auto"48    "Device to use for inference, e.g. 'cpu' or 'cuda', or 'mps'"49 50    backend: str = "torch"51    "Backend for inference, e.g. 'torch' (recommended for ROCm/Nvidia)"52    " or 'optimum' for onnx/tensorrt"53 54    model_warmup: bool = True55    "Warmup the model with the max batch size."56 57    engine: Any = None  #: :meta private:58    """Infinity's AsyncEmbeddingEngine."""59 60    # LLM call kwargs61    model_config = ConfigDict(62        extra="forbid",63        protected_namespaces=(),64    )65 66    @model_validator(mode="after")67    def validate_environment(self) -> Self:68        """Validate that api key and python package exists in environment."""69 70        try:71            from infinity_emb import AsyncEmbeddingEngine72        except ImportError:73            raise ImportError(74                "Please install the "75                "`pip install 'infinity_emb[optimum,torch]>=0.0.24'` "76                "package to use the InfinityEmbeddingsLocal."77            )78        self.engine = AsyncEmbeddingEngine(79            model_name_or_path=self.model,80            device=self.device,81            revision=self.revision,82            model_warmup=self.model_warmup,83            batch_size=self.batch_size,84            engine=self.backend,85        )86        return self87 88    async def __aenter__(self) -> None:89        """start the background worker.90        recommended usage is with the async with statement.91 92        async with InfinityEmbeddingsLocal(93            model="BAAI/bge-small-en-v1.5",94            revision=None,95            device="cpu",96        ) as embedder:97            embeddings = await engine.aembed_documents(["text1", "text2"])98        """99        await self.engine.__aenter__()100 101    async def __aexit__(self, *args: Any) -> None:102        """stop the background worker,103        required to free references to the pytorch model."""104        await self.engine.__aexit__(*args)105 106    async def aembed_documents(self, texts: List[str]) -> List[List[float]]:107        """Async call out to Infinity's embedding endpoint.108 109        Args:110            texts: The list of texts to embed.111 112        Returns:113            List of embeddings, one for each text.114        """115        if not self.engine.running:116            logger.warning(117                "Starting Infinity engine on the fly. This is not recommended."118                "Please start the engine before using it."119            )120            async with self:121                # spawning threadpool for multithreaded encode, tokenization122                embeddings, _ = await self.engine.embed(texts)123            # stopping threadpool on exit124            logger.warning("Stopped infinity engine after usage.")125        else:126            embeddings, _ = await self.engine.embed(texts)127        return embeddings128 129    async def aembed_query(self, text: str) -> List[float]:130        """Async call out to Infinity's embedding endpoint.131 132        Args:133            text: The text to embed.134 135        Returns:136            Embeddings for the text.137        """138        embeddings = await self.aembed_documents([text])139        return embeddings[0]140 141    def embed_documents(self, texts: List[str]) -> List[List[float]]:142        """143        This method is async only.144        """145        logger.warning(146            "This method is async only. "147            "Please use the async version `await aembed_documents`."148        )149        return asyncio.run(self.aembed_documents(texts))150 151    def embed_query(self, text: str) -> List[float]:152        """ """153        logger.warning(154            "This method is async only."155            " Please use the async version `await aembed_query`."156        )157        return asyncio.run(self.aembed_query(text))158 
codekingpro/portable-devtools · Team Ai