Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
google_palm.py104 linesDownload Raw Back to embeddings
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 
codekingpro/portable-devtools · Team Ai