codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4import warnings5from typing import (6 Any,7 Callable,8 Dict,9 List,10 Literal,11 Optional,12 Sequence,13 Set,14 Tuple,15 Union,16)17 18from langchain_core.embeddings import Embeddings19from langchain_core.utils import (20 get_from_dict_or_env,21 get_pydantic_field_names,22 pre_init,23)24from pydantic import BaseModel, ConfigDict, Field, model_validator25from tenacity import (26 AsyncRetrying,27 before_sleep_log,28 retry,29 retry_if_exception_type,30 stop_after_attempt,31 wait_exponential,32)33 34logger = logging.getLogger(__name__)35 36 37def _create_retry_decorator(embeddings: LocalAIEmbeddings) -> Callable[[Any], Any]:38 import openai39 40 min_seconds = 441 max_seconds = 1042 # Wait 2^x * 1 second between each retry starting with43 # 4 seconds, then up to 10 seconds, then 10 seconds afterwards44 return retry(45 reraise=True,46 stop=stop_after_attempt(embeddings.max_retries),47 wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),48 retry=(49 retry_if_exception_type(openai.error.Timeout)50 | retry_if_exception_type(openai.error.APIError)51 | retry_if_exception_type(openai.error.APIConnectionError)52 | retry_if_exception_type(openai.error.RateLimitError)53 | retry_if_exception_type(openai.error.ServiceUnavailableError)54 ),55 before_sleep=before_sleep_log(logger, logging.WARNING),56 )57 58 59def _async_retry_decorator(embeddings: LocalAIEmbeddings) -> Any:60 import openai61 62 min_seconds = 463 max_seconds = 1064 # Wait 2^x * 1 second between each retry starting with65 # 4 seconds, then up to 10 seconds, then 10 seconds afterwards66 async_retrying = AsyncRetrying(67 reraise=True,68 stop=stop_after_attempt(embeddings.max_retries),69 wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),70 retry=(71 retry_if_exception_type(openai.error.Timeout)72 | retry_if_exception_type(openai.error.APIError)73 | retry_if_exception_type(openai.error.APIConnectionError)74 | retry_if_exception_type(openai.error.RateLimitError)75 | retry_if_exception_type(openai.error.ServiceUnavailableError)76 ),77 before_sleep=before_sleep_log(logger, logging.WARNING),78 )79 80 def wrap(func: Callable) -> Callable:81 async def wrapped_f(*args: Any, **kwargs: Any) -> Callable:82 async for _ in async_retrying:83 return await func(*args, **kwargs)84 raise AssertionError("this is unreachable")85 86 return wrapped_f87 88 return wrap89 90 91# https://stackoverflow.com/questions/76469415/getting-embeddings-of-length-1-from-langchain-openaiembeddings92def _check_response(response: dict) -> dict:93 if any(len(d["embedding"]) == 1 for d in response["data"]):94 import openai95 96 raise openai.error.APIError("LocalAI API returned an empty embedding")97 return response98 99 100def embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:101 """Use tenacity to retry the embedding call."""102 retry_decorator = _create_retry_decorator(embeddings)103 104 @retry_decorator105 def _embed_with_retry(**kwargs: Any) -> Any:106 response = embeddings.client.create(**kwargs)107 return _check_response(response)108 109 return _embed_with_retry(**kwargs)110 111 112async def async_embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:113 """Use tenacity to retry the embedding call."""114 115 @_async_retry_decorator(embeddings)116 async def _async_embed_with_retry(**kwargs: Any) -> Any:117 response = await embeddings.client.acreate(**kwargs)118 return _check_response(response)119 120 return await _async_embed_with_retry(**kwargs)121 122 123class LocalAIEmbeddings(BaseModel, Embeddings):124 """LocalAI embedding models.125 126 Since LocalAI and OpenAI have 1:1 compatibility between APIs, this class127 uses the ``openai`` Python package's ``openai.Embedding`` as its client.128 Thus, you should have the ``openai`` python package installed, and defeat129 the environment variable ``OPENAI_API_KEY`` by setting to a random string.130 You also need to specify ``OPENAI_API_BASE`` to point to your LocalAI131 service endpoint.132 133 Example:134 .. code-block:: python135 136 from langchain_community.embeddings import LocalAIEmbeddings137 openai = LocalAIEmbeddings(138 openai_api_key="random-string",139 openai_api_base="http://localhost:8080"140 )141 142 """143 144 client: Any = None #: :meta private:145 model: str = "text-embedding-ada-002"146 deployment: str = model147 openai_api_version: Optional[str] = None148 openai_api_base: Optional[str] = None149 # to support explicit proxy for LocalAI150 openai_proxy: Optional[str] = None151 embedding_ctx_length: int = 8191152 """The maximum number of tokens to embed at once."""153 openai_api_key: Optional[str] = None154 openai_organization: Optional[str] = None155 allowed_special: Union[Literal["all"], Set[str]] = set()156 disallowed_special: Union[Literal["all"], Set[str], Sequence[str]] = "all"157 chunk_size: int = 1000158 """Maximum number of texts to embed in each batch"""159 max_retries: int = 6160 """Maximum number of retries to make when generating."""161 request_timeout: Optional[Union[float, Tuple[float, float]]] = None162 """Timeout in seconds for the LocalAI request."""163 headers: Any = None164 show_progress_bar: bool = False165 """Whether to show a progress bar when embedding."""166 model_kwargs: Dict[str, Any] = Field(default_factory=dict)167 """Holds any model parameters valid for `create` call not explicitly specified."""168 169 model_config = ConfigDict(extra="forbid", protected_namespaces=())170 171 @model_validator(mode="before")172 @classmethod173 def build_extra(cls, values: Dict[str, Any]) -> Any:174 """Build extra kwargs from additional params that were passed in."""175 all_required_field_names = get_pydantic_field_names(cls)176 extra = values.get("model_kwargs", {})177 for field_name in list(values):178 if field_name in extra:179 raise ValueError(f"Found {field_name} supplied twice.")180 if field_name not in all_required_field_names:181 warnings.warn(182 f"""WARNING! {field_name} is not default parameter.183 {field_name} was transferred to model_kwargs.184 Please confirm that {field_name} is what you intended."""185 )186 extra[field_name] = values.pop(field_name)187 188 invalid_model_kwargs = all_required_field_names.intersection(extra.keys())189 if invalid_model_kwargs:190 raise ValueError(191 f"Parameters {invalid_model_kwargs} should be specified explicitly. "192 f"Instead they were passed in as part of `model_kwargs` parameter."193 )194 195 values["model_kwargs"] = extra196 return values197 198 @pre_init199 def validate_environment(cls, values: Dict) -> Dict:200 """Validate that api key and python package exists in environment."""201 values["openai_api_key"] = get_from_dict_or_env(202 values, "openai_api_key", "OPENAI_API_KEY"203 )204 values["openai_api_base"] = get_from_dict_or_env(205 values,206 "openai_api_base",207 "OPENAI_API_BASE",208 default="",209 )210 values["openai_proxy"] = get_from_dict_or_env(211 values,212 "openai_proxy",213 "OPENAI_PROXY",214 default="",215 )216 217 default_api_version = ""218 values["openai_api_version"] = get_from_dict_or_env(219 values,220 "openai_api_version",221 "OPENAI_API_VERSION",222 default=default_api_version,223 )224 values["openai_organization"] = get_from_dict_or_env(225 values,226 "openai_organization",227 "OPENAI_ORGANIZATION",228 default="",229 )230 try:231 import openai232 233 values["client"] = openai.Embedding234 except ImportError:235 raise ImportError(236 "Could not import openai python package. "237 "Please install it with `pip install openai`."238 )239 return values240 241 @property242 def _invocation_params(self) -> Dict:243 openai_args = {244 "model": self.model,245 "request_timeout": self.request_timeout,246 "headers": self.headers,247 "api_key": self.openai_api_key,248 "organization": self.openai_organization,249 "api_base": self.openai_api_base,250 "api_version": self.openai_api_version,251 **self.model_kwargs,252 }253 if self.openai_proxy:254 import openai255 256 openai.proxy = {257 "http": self.openai_proxy,258 "https": self.openai_proxy,259 }260 return openai_args261 262 def _embedding_func(self, text: str, *, engine: str) -> List[float]:263 """Call out to LocalAI's embedding endpoint."""264 # handle large input text265 if self.model.endswith("001"):266 # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500267 # replace newlines, which can negatively affect performance.268 text = text.replace("\n", " ")269 return embed_with_retry(270 self,271 input=[text],272 **self._invocation_params,273 )["data"][0]["embedding"]274 275 async def _aembedding_func(self, text: str, *, engine: str) -> List[float]:276 """Call out to LocalAI's embedding endpoint."""277 # handle large input text278 if self.model.endswith("001"):279 # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500280 # replace newlines, which can negatively affect performance.281 text = text.replace("\n", " ")282 return (283 await async_embed_with_retry(284 self,285 input=[text],286 **self._invocation_params,287 )288 )["data"][0]["embedding"]289 290 def embed_documents(291 self, texts: List[str], chunk_size: Optional[int] = 0292 ) -> List[List[float]]:293 """Call out to LocalAI's embedding endpoint for embedding search docs.294 295 Args:296 texts: The list of texts to embed.297 chunk_size: The chunk size of embeddings. If None, will use the chunk size298 specified by the class.299 300 Returns:301 List of embeddings, one for each text.302 """303 # call _embedding_func for each text304 return [self._embedding_func(text, engine=self.deployment) for text in texts]305 306 async def aembed_documents(307 self, texts: List[str], chunk_size: Optional[int] = 0308 ) -> List[List[float]]:309 """Call out to LocalAI's embedding endpoint async for embedding search docs.310 311 Args:312 texts: The list of texts to embed.313 chunk_size: The chunk size of embeddings. If None, will use the chunk size314 specified by the class.315 316 Returns:317 List of embeddings, one for each text.318 """319 embeddings = []320 for text in texts:321 response = await self._aembedding_func(text, engine=self.deployment)322 embeddings.append(response)323 return embeddings324 325 def embed_query(self, text: str) -> List[float]:326 """Call out to LocalAI's embedding endpoint for embedding query text.327 328 Args:329 text: The text to embed.330 331 Returns:332 Embedding for the text.333 """334 embedding = self._embedding_func(text, engine=self.deployment)335 return embedding336 337 async def aembed_query(self, text: str) -> List[float]:338 """Call out to LocalAI's embedding endpoint async for embedding query text.339 340 Args:341 text: The text to embed.342 343 Returns:344 Embedding for the text.345 """346 embedding = await self._aembedding_func(text, engine=self.deployment)347 return embedding348 