Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
infinity.py325 linesDownload Raw Back to embeddings
1"""written under MIT Licence, Michael Feil 2023."""2 3import asyncio4from concurrent.futures import ThreadPoolExecutor5from typing import Any, Callable, Dict, List, Optional, Tuple6 7import aiohttp8import numpy as np9import requests10from langchain_core.embeddings import Embeddings11from langchain_core.utils import get_from_dict_or_env12from pydantic import BaseModel, ConfigDict, model_validator13 14__all__ = ["InfinityEmbeddings"]15 16 17class InfinityEmbeddings(BaseModel, Embeddings):18    """Self-hosted embedding models for `infinity` package.19 20    See https://github.com/michaelfeil/infinity21    This also works for text-embeddings-inference and other22    self-hosted openai-compatible servers.23 24    Infinity is a package to interact with Embedding Models on https://github.com/michaelfeil/infinity25 26 27    Example:28        .. code-block:: python29 30            from langchain_community.embeddings import InfinityEmbeddings31            InfinityEmbeddings(32                model="BAAI/bge-small",33                infinity_api_url="http://localhost:7997",34            )35    """36 37    model: str38    "Underlying Infinity model id."39 40    infinity_api_url: str = "http://localhost:7997"41    """Endpoint URL to use."""42 43    client: Any = None  #: :meta private:44    """Infinity client."""45 46    # LLM call kwargs47    model_config = ConfigDict(48        extra="forbid",49    )50 51    @model_validator(mode="before")52    @classmethod53    def validate_environment(cls, values: Dict) -> Any:54        """Validate that api key and python package exists in environment."""55 56        values["infinity_api_url"] = get_from_dict_or_env(57            values, "infinity_api_url", "INFINITY_API_URL"58        )59 60        values["client"] = TinyAsyncOpenAIInfinityEmbeddingClient(61            host=values["infinity_api_url"],62        )63        return values64 65    def embed_documents(self, texts: List[str]) -> List[List[float]]:66        """Call out to Infinity's embedding endpoint.67 68        Args:69            texts: The list of texts to embed.70 71        Returns:72            List of embeddings, one for each text.73        """74        embeddings = self.client.embed(75            model=self.model,76            texts=texts,77        )78        return embeddings79 80    async def aembed_documents(self, texts: List[str]) -> List[List[float]]:81        """Async call out to Infinity's embedding endpoint.82 83        Args:84            texts: The list of texts to embed.85 86        Returns:87            List of embeddings, one for each text.88        """89        embeddings = await self.client.aembed(90            model=self.model,91            texts=texts,92        )93        return embeddings94 95    def embed_query(self, text: str) -> List[float]:96        """Call out to Infinity's embedding endpoint.97 98        Args:99            text: The text to embed.100 101        Returns:102            Embeddings for the text.103        """104        return self.embed_documents([text])[0]105 106    async def aembed_query(self, text: str) -> List[float]:107        """Async call out to Infinity's embedding endpoint.108 109        Args:110            text: The text to embed.111 112        Returns:113            Embeddings for the text.114        """115        embeddings = await self.aembed_documents([text])116        return embeddings[0]117 118 119class TinyAsyncOpenAIInfinityEmbeddingClient:  #: :meta private:120    """Helper tool to embed Infinity.121 122    It is not a part of Langchain's stable API,123    direct use discouraged.124 125    Example:126        .. code-block:: python127 128 129            mini_client = TinyAsyncInfinityEmbeddingClient(130            )131            embeds = mini_client.embed(132                model="BAAI/bge-small",133                text=["doc1", "doc2"]134            )135            # or136            embeds = await mini_client.aembed(137                model="BAAI/bge-small",138                text=["doc1", "doc2"]139            )140 141    """142 143    def __init__(144        self,145        host: str = "http://localhost:7797/v1",146        aiosession: Optional[aiohttp.ClientSession] = None,147    ) -> None:148        self.host = host149        self.aiosession = aiosession150 151        if self.host is None or len(self.host) < 3:152            raise ValueError(" param `host` must be set to a valid url")153        self._batch_size = 128154 155    @staticmethod156    def _permute(157        texts: List[str], sorter: Callable = len158    ) -> Tuple[List[str], Callable]:159        """Sort texts in ascending order, and160        delivers a lambda expr, which can sort a same length list161        https://github.com/UKPLab/sentence-transformers/blob/162        c5f93f70eca933c78695c5bc686ceda59651ae3b/sentence_transformers/SentenceTransformer.py#L156163 164        Args:165            texts (List[str]): _description_166            sorter (Callable, optional): _description_. Defaults to len.167 168        Returns:169            Tuple[List[str], Callable]: _description_170 171        Example:172            ```173            texts = ["one","three","four"]174            perm_texts, undo = self._permute(texts)175            texts == undo(perm_texts)176            ```177        """178 179        if len(texts) == 1:180            # special case query181            return texts, lambda t: t182        length_sorted_idx = np.argsort([-sorter(sen) for sen in texts])183        texts_sorted = [texts[idx] for idx in length_sorted_idx]184 185        return texts_sorted, lambda unsorted_embeddings: [  # E731186            unsorted_embeddings[idx] for idx in np.argsort(length_sorted_idx)187        ]188 189    def _batch(self, texts: List[str]) -> List[List[str]]:190        """191        splits Lists of text parts into batches of size max `self._batch_size`192        When encoding vector database,193 194        Args:195            texts (List[str]): List of sentences196            self._batch_size (int, optional): max batch size of one request.197 198        Returns:199            List[List[str]]: Batches of List of sentences200        """201        if len(texts) == 1:202            # special case query203            return [texts]204        batches = []205        for start_index in range(0, len(texts), self._batch_size):206            batches.append(texts[start_index : start_index + self._batch_size])207        return batches208 209    @staticmethod210    def _unbatch(batch_of_texts: List[List[Any]]) -> List[Any]:211        if len(batch_of_texts) == 1 and len(batch_of_texts[0]) == 1:212            # special case query213            return batch_of_texts[0]214        texts = []215        for sublist in batch_of_texts:216            texts.extend(sublist)217        return texts218 219    def _kwargs_post_request(self, model: str, texts: List[str]) -> Dict[str, Any]:220        """Build the kwargs for the Post request, used by sync221 222        Args:223            model (str): _description_224            texts (List[str]): _description_225 226        Returns:227            Dict[str, Collection[str]]: _description_228        """229        return dict(230            url=f"{self.host}/embeddings",231            headers={232                # "accept": "application/json",233                "content-type": "application/json",234            },235            json=dict(236                input=texts,237                model=model,238            ),239        )240 241    def _sync_request_embed(242        self, model: str, batch_texts: List[str]243    ) -> List[List[float]]:244        response = requests.post(245            **self._kwargs_post_request(model=model, texts=batch_texts)246        )247        if response.status_code != 200:248            raise Exception(249                f"Infinity returned an unexpected response with status "250                f"{response.status_code}: {response.text}"251            )252        return [e["embedding"] for e in response.json()["data"]]253 254    def embed(self, model: str, texts: List[str]) -> List[List[float]]:255        """call the embedding of model256 257        Args:258            model (str): to embedding model259            texts (List[str]): List of sentences to embed.260 261        Returns:262            List[List[float]]: List of vectors for each sentence263        """264        perm_texts, unpermute_func = self._permute(texts)265        perm_texts_batched = self._batch(perm_texts)266 267        # Request268        map_args = (269            self._sync_request_embed,270            [model] * len(perm_texts_batched),271            perm_texts_batched,272        )273        if len(perm_texts_batched) == 1:274            embeddings_batch_perm = list(map(*map_args))275        else:276            with ThreadPoolExecutor(32) as p:277                embeddings_batch_perm = list(p.map(*map_args))278 279        embeddings_perm = self._unbatch(embeddings_batch_perm)280        embeddings = unpermute_func(embeddings_perm)281        return embeddings282 283    async def _async_request(284        self, session: aiohttp.ClientSession, kwargs: Dict[str, Any]285    ) -> List[List[float]]:286        async with session.post(**kwargs) as response:287            if response.status != 200:288                raise Exception(289                    f"Infinity returned an unexpected response with status "290                    f"{response.status}: {response.text}"291                )292            embedding = (await response.json())["data"]293            return [e["embedding"] for e in embedding]294 295    async def aembed(self, model: str, texts: List[str]) -> List[List[float]]:296        """call the embedding of model, async method297 298        Args:299            model (str): to embedding model300            texts (List[str]): List of sentences to embed.301 302        Returns:303            List[List[float]]: List of vectors for each sentence304        """305        perm_texts, unpermute_func = self._permute(texts)306        perm_texts_batched = self._batch(perm_texts)307 308        # Request309        async with aiohttp.ClientSession(310            trust_env=True, connector=aiohttp.TCPConnector(limit=32)311        ) as session:312            embeddings_batch_perm = await asyncio.gather(313                *[314                    self._async_request(315                        session=session,316                        kwargs=self._kwargs_post_request(model=model, texts=t),317                    )318                    for t in perm_texts_batched319                ]320            )321 322        embeddings_perm = self._unbatch(embeddings_batch_perm)323        embeddings = unpermute_func(embeddings_perm)324        return embeddings325 
codekingpro/portable-devtools · Team Ai