codekingpro/portable-devtools
114k
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 