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