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