codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4from typing import Any, Callable, Dict, List, Optional5 6from langchain_core.embeddings import Embeddings7from langchain_core.utils import get_from_dict_or_env, pre_init8from pydantic import BaseModel, ConfigDict9from tenacity import (10 before_sleep_log,11 retry,12 retry_if_exception_type,13 stop_after_attempt,14 wait_exponential,15)16 17logger = logging.getLogger(__name__)18 19 20def _create_retry_decorator() -> Callable[[Any], Any]:21 """Returns a tenacity retry decorator, preconfigured to handle PaLM exceptions"""22 import google.api_core.exceptions23 24 multiplier = 225 min_seconds = 126 max_seconds = 6027 max_retries = 1028 29 return retry(30 reraise=True,31 stop=stop_after_attempt(max_retries),32 wait=wait_exponential(multiplier=multiplier, min=min_seconds, max=max_seconds),33 retry=(34 retry_if_exception_type(google.api_core.exceptions.ResourceExhausted)35 | retry_if_exception_type(google.api_core.exceptions.ServiceUnavailable)36 | retry_if_exception_type(google.api_core.exceptions.GoogleAPIError)37 ),38 before_sleep=before_sleep_log(logger, logging.WARNING),39 )40 41 42def embed_with_retry(43 embeddings: GooglePalmEmbeddings, *args: Any, **kwargs: Any44) -> Any:45 """Use tenacity to retry the completion call."""46 retry_decorator = _create_retry_decorator()47 48 @retry_decorator49 def _embed_with_retry(*args: Any, **kwargs: Any) -> Any:50 return embeddings.client.generate_embeddings(*args, **kwargs)51 52 return _embed_with_retry(*args, **kwargs)53 54 55class GooglePalmEmbeddings(BaseModel, Embeddings):56 """Google's PaLM Embeddings APIs."""57 58 client: Any59 google_api_key: Optional[str]60 model_name: str = "models/embedding-gecko-001"61 """Model name to use."""62 show_progress_bar: bool = False63 """Whether to show a tqdm progress bar. Must have `tqdm` installed."""64 65 model_config = ConfigDict(protected_namespaces=())66 67 @pre_init68 def validate_environment(cls, values: Dict) -> Dict:69 """Validate api key, python package exists."""70 google_api_key = get_from_dict_or_env(71 values, "google_api_key", "GOOGLE_API_KEY"72 )73 try:74 import google.generativeai as genai75 76 genai.configure(api_key=google_api_key)77 except ImportError:78 raise ImportError("Could not import google.generativeai python package.")79 80 values["client"] = genai81 82 return values83 84 def embed_documents(self, texts: List[str]) -> List[List[float]]:85 if self.show_progress_bar:86 try:87 from tqdm import tqdm88 89 iter_ = tqdm(texts, desc="GooglePalmEmbeddings")90 except ImportError:91 logger.warning(92 "Unable to show progress bar because tqdm could not be imported. "93 "Please install with `pip install tqdm`."94 )95 iter_ = texts96 else:97 iter_ = texts98 return [self.embed_query(text) for text in iter_]99 100 def embed_query(self, text: str) -> List[float]:101 """Embed query text."""102 embedding = embed_with_retry(self, self.model_name, text)103 return embedding["embedding"]104 