Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
textembed.py351 linesDownload Raw Back to embeddings
1"""2TextEmbed: Embedding Inference Server3 4TextEmbed provides a high-throughput, low-latency solution for serving embeddings.5It supports various sentence-transformer models.6Now, it includes the ability to deploy image embedding models.7TextEmbed offers flexibility and scalability for diverse applications.8 9TextEmbed is maintained by Keval Dekivadiya and is licensed under the Apache-2.0 license.10"""  # noqa: E50111 12import asyncio13from concurrent.futures import ThreadPoolExecutor14from typing import Any, Callable, Dict, List, Optional, Tuple, Union15 16import aiohttp17import numpy as np18import requests19from langchain_core.embeddings import Embeddings20from langchain_core.utils import from_env, secret_from_env21from pydantic import BaseModel, ConfigDict, Field, SecretStr, model_validator22from typing_extensions import Self23 24__all__ = ["TextEmbedEmbeddings"]25 26 27class TextEmbedEmbeddings(BaseModel, Embeddings):28    """29    A class to handle embedding requests to the TextEmbed API.30 31    Attributes:32        model : The TextEmbed model ID to use for embeddings.33        api_url : The base URL for the TextEmbed API.34        api_key : The API key for authenticating with the TextEmbed API.35        client : The TextEmbed client instance.36 37    Example:38        .. code-block:: python39 40            from langchain_community.embeddings import TextEmbedEmbeddings41 42            embeddings = TextEmbedEmbeddings(43                model="sentence-transformers/clip-ViT-B-32",44                api_url="http://localhost:8000/v1",45                api_key="<API_KEY>"46            )47 48    For more information: https://github.com/kevaldekivadiya2415/textembed/blob/main/docs/setup.md49    """  # noqa: E50150 51    model: str52    """Underlying TextEmbed model id."""53 54    api_url: str = Field(55        default_factory=from_env(56            "TEXTEMBED_API_URL", default="http://localhost:8000/v1"57        )58    )59    """Endpoint URL to use."""60 61    api_key: SecretStr = Field(default_factory=secret_from_env("TEXTEMBED_API_KEY"))62    """API Key for authentication"""63 64    client: Any = None65    """TextEmbed client."""66 67    model_config = ConfigDict(68        extra="forbid",69    )70 71    @model_validator(mode="after")72    def validate_environment(self) -> Self:73        """Validate that api key and URL exist in the environment."""74        self.client = AsyncOpenAITextEmbedEmbeddingClient(75            host=self.api_url, api_key=self.api_key.get_secret_value()76        )77        return self78 79    def embed_documents(self, texts: List[str]) -> List[List[float]]:80        """Call out to TextEmbed's embedding endpoint.81 82        Args:83            texts (List[str]): The list of texts to embed.84 85        Returns:86            List[List[float]]: List of embeddings, one for each text.87        """88        embeddings = self.client.embed(89            model=self.model,90            texts=texts,91        )92        return embeddings93 94    async def aembed_documents(self, texts: List[str]) -> List[List[float]]:95        """Async call out to TextEmbed's embedding endpoint.96 97        Args:98            texts (List[str]): The list of texts to embed.99 100        Returns:101            List[List[float]]: List of embeddings, one for each text.102        """103        embeddings = await self.client.aembed(104            model=self.model,105            texts=texts,106        )107        return embeddings108 109    def embed_query(self, text: str) -> List[float]:110        """Call out to TextEmbed's embedding endpoint for a single query.111 112        Args:113            text (str): The text to embed.114 115        Returns:116            List[float]: Embeddings for the text.117        """118        return self.embed_documents([text])[0]119 120    async def aembed_query(self, text: str) -> List[float]:121        """Async call out to TextEmbed's embedding endpoint for a single query.122 123        Args:124            text (str): The text to embed.125 126        Returns:127            List[float]: Embeddings for the text.128        """129        embeddings = await self.aembed_documents([text])130        return embeddings[0]131 132 133class AsyncOpenAITextEmbedEmbeddingClient:134    """135    A client to handle synchronous and asynchronous requests to the TextEmbed API.136 137    Attributes:138        host (str): The base URL for the TextEmbed API.139        api_key (str): The API key for authenticating with the TextEmbed API.140        aiosession (Optional[aiohttp.ClientSession]): The aiohttp session for async requests.141        _batch_size (int): Maximum batch size for a single request.142    """  # noqa: E501143 144    def __init__(145        self,146        host: str = "http://localhost:8000/v1",147        api_key: Union[str, None] = None,148        aiosession: Optional[aiohttp.ClientSession] = None,149    ) -> None:150        self.host = host151        self.api_key = api_key152        self.aiosession = aiosession153 154        if self.host is None or len(self.host) < 3:155            raise ValueError("Parameter `host` must be set to a valid URL")156        self._batch_size = 256157 158    @staticmethod159    def _permute(160        texts: List[str], sorter: Callable = len161    ) -> Tuple[List[str], Callable]:162        """163        Sorts texts in ascending order and provides a function to restore the original order.164 165        Args:166            texts (List[str]): List of texts to sort.167            sorter (Callable, optional): Sorting function, defaults to length.168 169        Returns:170            Tuple[List[str], Callable]: Sorted texts and a function to restore original order.171        """  # noqa: E501172        if len(texts) == 1:173            return texts, lambda t: t174        length_sorted_idx = np.argsort([-sorter(sen) for sen in texts])175        texts_sorted = [texts[idx] for idx in length_sorted_idx]176 177        return texts_sorted, lambda unsorted_embeddings: [178            unsorted_embeddings[idx] for idx in np.argsort(length_sorted_idx)179        ]180 181    def _batch(self, texts: List[str]) -> List[List[str]]:182        """183        Splits a list of texts into batches of size max `self._batch_size`.184 185        Args:186            texts (List[str]): List of texts to split.187 188        Returns:189            List[List[str]]: List of batches of texts.190        """191        if len(texts) == 1:192            return [texts]193        batches = []194        for start_index in range(0, len(texts), self._batch_size):195            batches.append(texts[start_index : start_index + self._batch_size])196        return batches197 198    @staticmethod199    def _unbatch(batch_of_texts: List[List[Any]]) -> List[Any]:200        """201        Merges batches of texts into a single list.202 203        Args:204            batch_of_texts (List[List[Any]]): List of batches of texts.205 206        Returns:207            List[Any]: Merged list of texts.208        """209        if len(batch_of_texts) == 1 and len(batch_of_texts[0]) == 1:210            return batch_of_texts[0]211        texts = []212        for sublist in batch_of_texts:213            texts.extend(sublist)214        return texts215 216    def _kwargs_post_request(self, model: str, texts: List[str]) -> Dict[str, Any]:217        """218        Builds the kwargs for the POST request, used by sync method.219 220        Args:221            model (str): The model to use for embedding.222            texts (List[str]): List of texts to embed.223 224        Returns:225            Dict[str, Any]: Dictionary of POST request parameters.226        """227        return dict(228            url=f"{self.host}/embedding",229            headers={230                "accept": "application/json",231                "content-type": "application/json",232                "Authorization": f"Bearer {self.api_key}",233            },234            json=dict(235                input=texts,236                model=model,237            ),238        )239 240    def _sync_request_embed(241        self, model: str, batch_texts: List[str]242    ) -> List[List[float]]:243        """244        Sends a synchronous request to the embedding endpoint.245 246        Args:247            model (str): The model to use for embedding.248            batch_texts (List[str]): Batch of texts to embed.249 250        Returns:251            List[List[float]]: List of embeddings for the batch.252 253        Raises:254            Exception: If the response status is not 200.255        """256        response = requests.post(257            **self._kwargs_post_request(model=model, texts=batch_texts)258        )259        if response.status_code != 200:260            raise Exception(261                f"TextEmbed responded with an unexpected status message "262                f"{response.status_code}: {response.text}"263            )264        return [e["embedding"] for e in response.json()["data"]]265 266    def embed(self, model: str, texts: List[str]) -> List[List[float]]:267        """268        Embeds a list of texts synchronously.269 270        Args:271            model (str): The model to use for embedding.272            texts (List[str]): List of texts to embed.273 274        Returns:275            List[List[float]]: List of embeddings for the texts.276        """277        perm_texts, unpermute_func = self._permute(texts)278        perm_texts_batched = self._batch(perm_texts)279 280        # Request281        map_args = (282            self._sync_request_embed,283            [model] * len(perm_texts_batched),284            perm_texts_batched,285        )286        if len(perm_texts_batched) == 1:287            embeddings_batch_perm = list(map(*map_args))288        else:289            with ThreadPoolExecutor(32) as p:290                embeddings_batch_perm = list(p.map(*map_args))291 292        embeddings_perm = self._unbatch(embeddings_batch_perm)293        embeddings = unpermute_func(embeddings_perm)294        return embeddings295 296    async def _async_request(297        self, session: aiohttp.ClientSession, **kwargs: Dict[str, Any]298    ) -> List[List[float]]:299        """300        Sends an asynchronous request to the embedding endpoint.301 302        Args:303            session (aiohttp.ClientSession): The aiohttp session for the request.304            kwargs (Dict[str, Any]): Dictionary of POST request parameters.305 306        Returns:307            List[List[float]]: List of embeddings for the request.308 309        Raises:310            Exception: If the response status is not 200.311        """312        async with session.post(**kwargs) as response:  # type: ignore[arg-type]313            if response.status != 200:314                raise Exception(315                    f"TextEmbed responded with an unexpected status message "316                    f"{response.status}: {response.text}"317                )318            embedding = (await response.json())["data"]319            return [e["embedding"] for e in embedding]320 321    async def aembed(self, model: str, texts: List[str]) -> List[List[float]]:322        """323        Embeds a list of texts asynchronously.324 325        Args:326            model (str): The model to use for embedding.327            texts (List[str]): List of texts to embed.328 329        Returns:330            List[List[float]]: List of embeddings for the texts.331        """332        perm_texts, unpermute_func = self._permute(texts)333        perm_texts_batched = self._batch(perm_texts)334 335        async with aiohttp.ClientSession(336            connector=aiohttp.TCPConnector(limit=32)337        ) as session:338            embeddings_batch_perm = await asyncio.gather(339                *[340                    self._async_request(341                        session=session,342                        **self._kwargs_post_request(model=model, texts=t),343                    )344                    for t in perm_texts_batched345                ]346            )347 348        embeddings_perm = self._unbatch(embeddings_batch_perm)349        embeddings = unpermute_func(embeddings_perm)350        return embeddings351 
codekingpro/portable-devtools · Team Ai