Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
vertexai.py362 linesDownload Raw Back to embeddings
1import logging2import re3import string4import threading5from concurrent.futures import ThreadPoolExecutor, wait6from typing import Any, Dict, List, Literal, Optional, Tuple7 8from langchain_core._api.deprecation import deprecated9from langchain_core.embeddings import Embeddings10from langchain_core.language_models.llms import create_base_retry_decorator11from langchain_core.utils import pre_init12 13from langchain_community.llms.vertexai import _VertexAICommon14from langchain_community.utilities.vertexai import raise_vertex_import_error15 16logger = logging.getLogger(__name__)17 18_MAX_TOKENS_PER_BATCH = 2000019_MAX_BATCH_SIZE = 25020_MIN_BATCH_SIZE = 521 22 23@deprecated(24    since="0.0.12",25    removal="1.0",26    alternative_import="langchain_google_vertexai.VertexAIEmbeddings",27)28class VertexAIEmbeddings(_VertexAICommon, Embeddings):29    """Google Cloud VertexAI embedding models."""30 31    # Instance context32    instance: Dict[str, Any] = {}  #: :meta private:33    show_progress_bar: bool = False34    """Whether to show a tqdm progress bar. Must have `tqdm` installed."""35 36    @pre_init37    def validate_environment(cls, values: Dict) -> Dict:38        """Validates that the python package exists in environment."""39        cls._try_init_vertexai(values)40        if values["model_name"] == "textembedding-gecko-default":41            logger.warning(42                "Model_name will become a required arg for VertexAIEmbeddings "43                "starting from Feb-01-2024. Currently the default is set to "44                "textembedding-gecko@001"45            )46            values["model_name"] = "textembedding-gecko@001"47        try:48            from vertexai.language_models import TextEmbeddingModel49        except ImportError:50            raise_vertex_import_error()51        values["client"] = TextEmbeddingModel.from_pretrained(values["model_name"])52        return values53 54    def __init__(55        self,56        # the default value would be removed after Feb-01-202457        model_name: str = "textembedding-gecko-default",58        project: Optional[str] = None,59        location: str = "us-central1",60        request_parallelism: int = 5,61        max_retries: int = 6,62        credentials: Optional[Any] = None,63        **kwargs: Any,64    ):65        """Initialize the sentence_transformer."""66        super().__init__(67            project=project,68            location=location,69            credentials=credentials,70            request_parallelism=request_parallelism,71            max_retries=max_retries,72            model_name=model_name,73            **kwargs,74        )75        self.instance["max_batch_size"] = kwargs.get("max_batch_size", _MAX_BATCH_SIZE)76        self.instance["batch_size"] = self.instance["max_batch_size"]77        self.instance["min_batch_size"] = kwargs.get("min_batch_size", _MIN_BATCH_SIZE)78        self.instance["min_good_batch_size"] = self.instance["min_batch_size"]79        self.instance["lock"] = threading.Lock()80        self.instance["batch_size_validated"] = False81        self.instance["task_executor"] = ThreadPoolExecutor(82            max_workers=request_parallelism83        )84        self.instance[85            "embeddings_task_type_supported"86        ] = not self.client._endpoint_name.endswith("/textembedding-gecko@001")87 88    @staticmethod89    def _split_by_punctuation(text: str) -> List[str]:90        """Splits a string by punctuation and whitespace characters."""91        split_by = string.punctuation + "\t\n "92        pattern = f"([{split_by}])"93        # Using re.split to split the text based on the pattern94        return [segment for segment in re.split(pattern, text) if segment]95 96    @staticmethod97    def _prepare_batches(texts: List[str], batch_size: int) -> List[List[str]]:98        """Splits texts in batches based on current maximum batch size99        and maximum tokens per request.100        """101        text_index = 0102        texts_len = len(texts)103        batch_token_len = 0104        batches: List[List[str]] = []105        current_batch: List[str] = []106        if texts_len == 0:107            return []108        while text_index < texts_len:109            current_text = texts[text_index]110            # Number of tokens per a text is conservatively estimated111            # as 2 times number of words, punctuation and whitespace characters.112            # Using `count_tokens` API will make batching too expensive.113            # Utilizing a tokenizer, would add a dependency that would not114            # necessarily be reused by the application using this class.115            current_text_token_cnt = (116                len(VertexAIEmbeddings._split_by_punctuation(current_text)) * 2117            )118            end_of_batch = False119            if current_text_token_cnt > _MAX_TOKENS_PER_BATCH:120                # Current text is too big even for a single batch.121                # Such request will fail, but we still make a batch122                # so that the app can get the error from the API.123                if len(current_batch) > 0:124                    # Adding current batch if not empty.125                    batches.append(current_batch)126                current_batch = [current_text]127                text_index += 1128                end_of_batch = True129            elif (130                batch_token_len + current_text_token_cnt > _MAX_TOKENS_PER_BATCH131                or len(current_batch) == batch_size132            ):133                end_of_batch = True134            else:135                if text_index == texts_len - 1:136                    # Last element - even though the batch may be not big,137                    # we still need to make it.138                    end_of_batch = True139                batch_token_len += current_text_token_cnt140                current_batch.append(current_text)141                text_index += 1142            if end_of_batch:143                batches.append(current_batch)144                current_batch = []145                batch_token_len = 0146        return batches147 148    def _get_embeddings_with_retry(149        self, texts: List[str], embeddings_type: Optional[str] = None150    ) -> List[List[float]]:151        """Makes a Vertex AI model request with retry logic."""152        from google.api_core.exceptions import (153            Aborted,154            DeadlineExceeded,155            ResourceExhausted,156            ServiceUnavailable,157        )158 159        errors = [160            ResourceExhausted,161            ServiceUnavailable,162            Aborted,163            DeadlineExceeded,164        ]165        retry_decorator = create_base_retry_decorator(166            error_types=errors,167            max_retries=self.max_retries,168        )169 170        @retry_decorator171        def _completion_with_retry(texts_to_process: List[str]) -> Any:172            if embeddings_type and self.instance["embeddings_task_type_supported"]:173                from vertexai.language_models import TextEmbeddingInput174 175                requests = [176                    TextEmbeddingInput(text=t, task_type=embeddings_type)177                    for t in texts_to_process178                ]179            else:180                requests = texts_to_process181            embeddings = self.client.get_embeddings(requests)182            return [embs.values for embs in embeddings]183 184        return _completion_with_retry(texts)185 186    def _prepare_and_validate_batches(187        self, texts: List[str], embeddings_type: Optional[str] = None188    ) -> Tuple[List[List[float]], List[List[str]]]:189        """Prepares text batches with one-time validation of batch size.190        Batch size varies between GCP regions and individual project quotas.191        # Returns embeddings of the first text batch that went through,192        # and text batches for the rest of the texts.193        """194        from google.api_core.exceptions import InvalidArgument195 196        batches = VertexAIEmbeddings._prepare_batches(197            texts, self.instance["batch_size"]198        )199        # If batch size if less or equal to one that went through before,200        # then keep batches as they are.201        if len(batches[0]) <= self.instance["min_good_batch_size"]:202            return [], batches203        with self.instance["lock"]:204            # If largest possible batch size was validated205            # while waiting for the lock, then check for rebuilding206            # our batches, and return.207            if self.instance["batch_size_validated"]:208                if len(batches[0]) <= self.instance["batch_size"]:209                    return [], batches210                else:211                    return [], VertexAIEmbeddings._prepare_batches(212                        texts, self.instance["batch_size"]213                    )214            # Figure out largest possible batch size by trying to push215            # batches and lowering their size in half after every failure.216            first_batch = batches[0]217            first_result = []218            had_failure = False219            while True:220                try:221                    first_result = self._get_embeddings_with_retry(222                        first_batch, embeddings_type223                    )224                    break225                except InvalidArgument:226                    had_failure = True227                    first_batch_len = len(first_batch)228                    if first_batch_len == self.instance["min_batch_size"]:229                        raise230                    first_batch_len = max(231                        self.instance["min_batch_size"], int(first_batch_len / 2)232                    )233                    first_batch = first_batch[:first_batch_len]234            first_batch_len = len(first_batch)235            self.instance["min_good_batch_size"] = max(236                self.instance["min_good_batch_size"], first_batch_len237            )238            # If had a failure and recovered239            # or went through with the max size, then it's a legit batch size.240            if had_failure or first_batch_len == self.instance["max_batch_size"]:241                self.instance["batch_size"] = first_batch_len242                self.instance["batch_size_validated"] = True243                # If batch size was updated,244                # rebuild batches with the new batch size245                # (texts that went through are excluded here).246                if first_batch_len != self.instance["max_batch_size"]:247                    batches = VertexAIEmbeddings._prepare_batches(248                        texts[first_batch_len:], self.instance["batch_size"]249                    )250            else:251                # Still figuring out max batch size.252                batches = batches[1:]253        # Returning embeddings of the first text batch that went through,254        # and text batches for the rest of texts.255        return first_result, batches256 257    def embed(258        self,259        texts: List[str],260        batch_size: int = 0,261        embeddings_task_type: Optional[262            Literal[263                "RETRIEVAL_QUERY",264                "RETRIEVAL_DOCUMENT",265                "SEMANTIC_SIMILARITY",266                "CLASSIFICATION",267                "CLUSTERING",268            ]269        ] = None,270    ) -> List[List[float]]:271        """Embed a list of strings.272 273        Args:274            texts: List[str] The list of strings to embed.275            batch_size: [int] The batch size of embeddings to send to the model.276                If zero, then the largest batch size will be detected dynamically277                at the first request, starting from 250, down to 5.278            embeddings_task_type: [str] optional embeddings task type,279                one of the following280                    RETRIEVAL_QUERY	- Text is a query281                                      in a search/retrieval setting.282                    RETRIEVAL_DOCUMENT - Text is a document283                                         in a search/retrieval setting.284                    SEMANTIC_SIMILARITY - Embeddings will be used285                                          for Semantic Textual Similarity (STS).286                    CLASSIFICATION - Embeddings will be used for classification.287                    CLUSTERING - Embeddings will be used for clustering.288 289        Returns:290            List of embeddings, one for each text.291        """292        if len(texts) == 0:293            return []294        embeddings: List[List[float]] = []295        first_batch_result: List[List[float]] = []296        if batch_size > 0:297            # Fixed batch size.298            batches = VertexAIEmbeddings._prepare_batches(texts, batch_size)299        else:300            # Dynamic batch size, starting from 250 at the first call.301            first_batch_result, batches = self._prepare_and_validate_batches(302                texts, embeddings_task_type303            )304        # First batch result may have some embeddings already.305        # In such case, batches have texts that were not processed yet.306        embeddings.extend(first_batch_result)307        tasks = []308        if self.show_progress_bar:309            try:310                from tqdm import tqdm311 312                iter_ = tqdm(batches, desc="VertexAIEmbeddings")313            except ImportError:314                logger.warning(315                    "Unable to show progress bar because tqdm could not be imported. "316                    "Please install with `pip install tqdm`."317                )318                iter_ = batches319        else:320            iter_ = batches321        for batch in iter_:322            tasks.append(323                self.instance["task_executor"].submit(324                    self._get_embeddings_with_retry,325                    texts=batch,326                    embeddings_type=embeddings_task_type,327                )328            )329        if len(tasks) > 0:330            wait(tasks)331        for t in tasks:332            embeddings.extend(t.result())333        return embeddings334 335    def embed_documents(336        self, texts: List[str], batch_size: int = 0337    ) -> List[List[float]]:338        """Embed a list of documents.339 340        Args:341            texts: List[str] The list of texts to embed.342            batch_size: [int] The batch size of embeddings to send to the model.343                If zero, then the largest batch size will be detected dynamically344                at the first request, starting from 250, down to 5.345 346        Returns:347            List of embeddings, one for each text.348        """349        return self.embed(texts, batch_size, "RETRIEVAL_DOCUMENT")350 351    def embed_query(self, text: str) -> List[float]:352        """Embed a text.353 354        Args:355            text: The text to embed.356 357        Returns:358            Embedding for the text.359        """360        embeddings = self.embed([text], 1, "RETRIEVAL_QUERY")361        return embeddings[0]362 
codekingpro/portable-devtools · Team Ai