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