codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from concurrent.futures import Executor, ThreadPoolExecutor4from typing import TYPE_CHECKING, Any, ClassVar, Dict, Iterator, List, Optional, Union5 6from langchain_core._api.deprecation import deprecated7from langchain_core.callbacks.manager import (8 AsyncCallbackManagerForLLMRun,9 CallbackManagerForLLMRun,10)11from langchain_core.language_models.llms import BaseLLM12from langchain_core.outputs import Generation, GenerationChunk, LLMResult13from langchain_core.utils import pre_init14from pydantic import BaseModel, ConfigDict, Field15 16from langchain_community.utilities.vertexai import (17 create_retry_decorator,18 get_client_info,19 init_vertexai,20 raise_vertex_import_error,21)22 23if TYPE_CHECKING:24 from google.cloud.aiplatform.gapic import (25 PredictionServiceAsyncClient,26 PredictionServiceClient,27 )28 from google.cloud.aiplatform.models import Prediction29 from google.protobuf.struct_pb2 import Value30 from vertexai.language_models._language_models import (31 TextGenerationResponse,32 _LanguageModel,33 )34 from vertexai.preview.generative_models import Image35 36# This is for backwards compatibility37# We can remove after `langchain` stops importing it38_response_to_generation = None39stream_completion_with_retry = None40 41 42def is_codey_model(model_name: str) -> bool:43 """Return True if the model name is a Codey model."""44 return "code" in model_name45 46 47def is_gemini_model(model_name: str) -> bool:48 """Return True if the model name is a Gemini model."""49 return model_name is not None and "gemini" in model_name50 51 52def completion_with_retry(53 llm: VertexAI,54 prompt: List[Union[str, "Image"]],55 stream: bool = False,56 is_gemini: bool = False,57 run_manager: Optional[CallbackManagerForLLMRun] = None,58 **kwargs: Any,59) -> Any:60 """Use tenacity to retry the completion call."""61 retry_decorator = create_retry_decorator(llm, run_manager=run_manager)62 63 @retry_decorator64 def _completion_with_retry(65 prompt: List[Union[str, "Image"]], is_gemini: bool = False, **kwargs: Any66 ) -> Any:67 if is_gemini:68 return llm.client.generate_content(69 prompt, stream=stream, generation_config=kwargs70 )71 else:72 if stream:73 return llm.client.predict_streaming(prompt[0], **kwargs)74 return llm.client.predict(prompt[0], **kwargs)75 76 return _completion_with_retry(prompt, is_gemini, **kwargs)77 78 79async def acompletion_with_retry(80 llm: VertexAI,81 prompt: str,82 is_gemini: bool = False,83 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,84 **kwargs: Any,85) -> Any:86 """Use tenacity to retry the completion call."""87 retry_decorator = create_retry_decorator(llm, run_manager=run_manager)88 89 @retry_decorator90 async def _acompletion_with_retry(91 prompt: str, is_gemini: bool = False, **kwargs: Any92 ) -> Any:93 if is_gemini:94 return await llm.client.generate_content_async(95 prompt, generation_config=kwargs96 )97 return await llm.client.predict_async(prompt, **kwargs)98 99 return await _acompletion_with_retry(prompt, is_gemini, **kwargs)100 101 102class _VertexAIBase(BaseModel):103 model_config = ConfigDict(protected_namespaces=())104 105 project: Optional[str] = None106 "The default GCP project to use when making Vertex API calls."107 location: str = "us-central1"108 "The default location to use when making API calls."109 request_parallelism: int = 5110 "The amount of parallelism allowed for requests issued to VertexAI models. "111 "Default is 5."112 max_retries: int = 6113 """The maximum number of retries to make when generating."""114 task_executor: ClassVar[Optional[Executor]] = Field(default=None, exclude=True)115 stop: Optional[List[str]] = None116 "Optional list of stop words to use when generating."117 model_name: Optional[str] = None118 "Underlying model name."119 120 @classmethod121 def _get_task_executor(cls, request_parallelism: int = 5) -> Executor:122 if cls.task_executor is None:123 cls.task_executor = ThreadPoolExecutor(max_workers=request_parallelism)124 return cls.task_executor125 126 127class _VertexAICommon(_VertexAIBase):128 client: "_LanguageModel" = None #: :meta private:129 client_preview: "_LanguageModel" = None #: :meta private:130 model_name: str131 "Underlying model name."132 temperature: float = 0.0133 "Sampling temperature, it controls the degree of randomness in token selection."134 max_output_tokens: int = 128135 "Token limit determines the maximum amount of text output from one prompt."136 top_p: float = 0.95137 "Tokens are selected from most probable to least until the sum of their "138 "probabilities equals the top-p value. Top-p is ignored for Codey models."139 top_k: int = 40140 "How the model selects tokens for output, the next token is selected from "141 "among the top-k most probable tokens. Top-k is ignored for Codey models."142 credentials: Any = Field(default=None, exclude=True)143 "The default custom credentials (google.auth.credentials.Credentials) to use "144 "when making API calls. If not provided, credentials will be ascertained from "145 "the environment."146 n: int = 1147 """How many completions to generate for each prompt."""148 streaming: bool = False149 """Whether to stream the results or not."""150 151 @property152 def _llm_type(self) -> str:153 return "vertexai"154 155 @property156 def is_codey_model(self) -> bool:157 return is_codey_model(self.model_name)158 159 @property160 def _is_gemini_model(self) -> bool:161 return is_gemini_model(self.model_name)162 163 @property164 def _identifying_params(self) -> Dict[str, Any]:165 """Gets the identifying parameters."""166 return {**{"model_name": self.model_name}, **self._default_params}167 168 @property169 def _default_params(self) -> Dict[str, Any]:170 params = {171 "temperature": self.temperature,172 "max_output_tokens": self.max_output_tokens,173 "candidate_count": self.n,174 }175 if not self.is_codey_model:176 params.update(177 {178 "top_k": self.top_k,179 "top_p": self.top_p,180 }181 )182 return params183 184 @classmethod185 def _try_init_vertexai(cls, values: Dict) -> None:186 allowed_params = ["project", "location", "credentials"]187 params = {k: v for k, v in values.items() if k in allowed_params}188 init_vertexai(**params)189 return None190 191 def _prepare_params(192 self,193 stop: Optional[List[str]] = None,194 stream: bool = False,195 **kwargs: Any,196 ) -> dict:197 stop_sequences = stop or self.stop198 params_mapping = {"n": "candidate_count"}199 params = {params_mapping.get(k, k): v for k, v in kwargs.items()}200 params = {**self._default_params, "stop_sequences": stop_sequences, **params}201 if stream or self.streaming:202 params.pop("candidate_count")203 return params204 205 206@deprecated(207 since="0.0.12",208 removal="1.0",209 alternative_import="langchain_google_vertexai.VertexAI",210)211class VertexAI(_VertexAICommon, BaseLLM):212 """Google Vertex AI large language models."""213 214 model_name: str = "text-bison"215 "The name of the Vertex AI large language model."216 tuned_model_name: Optional[str] = None217 "The name of a tuned model. If provided, model_name is ignored."218 219 @classmethod220 def is_lc_serializable(self) -> bool:221 return True222 223 @classmethod224 def get_lc_namespace(cls) -> List[str]:225 """Get the namespace of the langchain object."""226 return ["langchain", "llms", "vertexai"]227 228 @pre_init229 def validate_environment(cls, values: Dict) -> Dict:230 """Validate that the python package exists in environment."""231 tuned_model_name = values.get("tuned_model_name")232 model_name = values["model_name"]233 is_gemini = is_gemini_model(values["model_name"])234 cls._try_init_vertexai(values)235 try:236 from vertexai.language_models import (237 CodeGenerationModel,238 TextGenerationModel,239 )240 from vertexai.preview.language_models import (241 CodeGenerationModel as PreviewCodeGenerationModel,242 )243 from vertexai.preview.language_models import (244 TextGenerationModel as PreviewTextGenerationModel,245 )246 247 if is_gemini:248 from vertexai.preview.generative_models import (249 GenerativeModel,250 )251 252 if is_codey_model(model_name):253 model_cls = CodeGenerationModel254 preview_model_cls = PreviewCodeGenerationModel255 elif is_gemini:256 model_cls = GenerativeModel257 preview_model_cls = GenerativeModel258 else:259 model_cls = TextGenerationModel260 preview_model_cls = PreviewTextGenerationModel261 262 if tuned_model_name:263 values["client"] = model_cls.get_tuned_model(tuned_model_name)264 values["client_preview"] = preview_model_cls.get_tuned_model(265 tuned_model_name266 )267 else:268 if is_gemini:269 values["client"] = model_cls(model_name=model_name)270 values["client_preview"] = preview_model_cls(model_name=model_name)271 else:272 values["client"] = model_cls.from_pretrained(model_name)273 values["client_preview"] = preview_model_cls.from_pretrained(274 model_name275 )276 277 except ImportError:278 raise_vertex_import_error()279 280 if values["streaming"] and values["n"] > 1:281 raise ValueError("Only one candidate can be generated with streaming!")282 return values283 284 def get_num_tokens(self, text: str) -> int:285 """Get the number of tokens present in the text.286 287 Useful for checking if an input will fit in a model's context window.288 289 Args:290 text: The string input to tokenize.291 292 Returns:293 The integer number of tokens in the text.294 """295 try:296 result = self.client_preview.count_tokens([text])297 except AttributeError:298 raise_vertex_import_error()299 300 return result.total_tokens301 302 def _response_to_generation(303 self, response: TextGenerationResponse304 ) -> GenerationChunk:305 """Converts a stream response to a generation chunk."""306 try:307 generation_info = {308 "is_blocked": response.is_blocked,309 "safety_attributes": response.safety_attributes,310 }311 except Exception:312 generation_info = None313 return GenerationChunk(text=response.text, generation_info=generation_info)314 315 def _generate(316 self,317 prompts: List[str],318 stop: Optional[List[str]] = None,319 run_manager: Optional[CallbackManagerForLLMRun] = None,320 stream: Optional[bool] = None,321 **kwargs: Any,322 ) -> LLMResult:323 should_stream = stream if stream is not None else self.streaming324 params = self._prepare_params(stop=stop, stream=should_stream, **kwargs)325 generations: List[List[Generation]] = []326 for prompt in prompts:327 if should_stream:328 generation = GenerationChunk(text="")329 for chunk in self._stream(330 prompt, stop=stop, run_manager=run_manager, **kwargs331 ):332 generation += chunk333 generations.append([generation])334 else:335 res = completion_with_retry(336 self,337 [prompt],338 stream=should_stream,339 is_gemini=self._is_gemini_model,340 run_manager=run_manager,341 **params,342 )343 generations.append(344 [self._response_to_generation(r) for r in res.candidates]345 )346 return LLMResult(generations=generations)347 348 async def _agenerate(349 self,350 prompts: List[str],351 stop: Optional[List[str]] = None,352 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,353 **kwargs: Any,354 ) -> LLMResult:355 params = self._prepare_params(stop=stop, **kwargs)356 generations = []357 for prompt in prompts:358 res = await acompletion_with_retry(359 self,360 prompt,361 is_gemini=self._is_gemini_model,362 run_manager=run_manager,363 **params,364 )365 generations.append(366 [self._response_to_generation(r) for r in res.candidates]367 )368 return LLMResult(generations=generations) # type: ignore[arg-type]369 370 def _stream(371 self,372 prompt: str,373 stop: Optional[List[str]] = None,374 run_manager: Optional[CallbackManagerForLLMRun] = None,375 **kwargs: Any,376 ) -> Iterator[GenerationChunk]:377 params = self._prepare_params(stop=stop, stream=True, **kwargs)378 for stream_resp in completion_with_retry(379 self,380 [prompt],381 stream=True,382 is_gemini=self._is_gemini_model,383 run_manager=run_manager,384 **params,385 ):386 chunk = self._response_to_generation(stream_resp)387 if run_manager:388 run_manager.on_llm_new_token(389 chunk.text,390 chunk=chunk,391 verbose=self.verbose,392 )393 yield chunk394 395 396@deprecated(397 since="0.0.12",398 removal="1.0",399 alternative_import="langchain_google_vertexai.VertexAIModelGarden",400)401class VertexAIModelGarden(_VertexAIBase, BaseLLM):402 """Vertex AI Model Garden large language models."""403 404 client: "PredictionServiceClient" = (405 None #: :meta private: # type: ignore[assignment]406 )407 async_client: "PredictionServiceAsyncClient" = (408 None #: :meta private: # type: ignore[assignment]409 )410 endpoint_id: str411 "A name of an endpoint where the model has been deployed."412 allowed_model_args: Optional[List[str]] = None413 "Allowed optional args to be passed to the model."414 prompt_arg: str = "prompt"415 result_arg: Optional[str] = "generated_text"416 "Set result_arg to None if output of the model is expected to be a string."417 "Otherwise, if it's a dict, provided an argument that contains the result."418 419 @pre_init420 def validate_environment(cls, values: Dict) -> Dict:421 """Validate that the python package exists in environment."""422 try:423 from google.api_core.client_options import ClientOptions424 from google.cloud.aiplatform.gapic import (425 PredictionServiceAsyncClient,426 PredictionServiceClient,427 )428 except ImportError:429 raise_vertex_import_error()430 431 if not values["project"]:432 raise ValueError(433 "A GCP project should be provided to run inference on Model Garden!"434 )435 436 client_options = ClientOptions(437 api_endpoint=f"{values['location']}-aiplatform.googleapis.com"438 )439 client_info = get_client_info(module="vertex-ai-model-garden")440 values["client"] = PredictionServiceClient(441 client_options=client_options, client_info=client_info442 )443 values["async_client"] = PredictionServiceAsyncClient(444 client_options=client_options, client_info=client_info445 )446 return values447 448 @property449 def endpoint_path(self) -> str:450 return self.client.endpoint_path(451 project=self.project,452 location=self.location,453 endpoint=self.endpoint_id,454 )455 456 @property457 def _llm_type(self) -> str:458 return "vertexai_model_garden"459 460 def _prepare_request(self, prompts: List[str], **kwargs: Any) -> List["Value"]:461 try:462 from google.protobuf import json_format463 from google.protobuf.struct_pb2 import Value464 except ImportError:465 raise ImportError(466 "protobuf package not found, please install it with"467 " `pip install protobuf`"468 )469 instances = []470 for prompt in prompts:471 if self.allowed_model_args:472 instance = {473 k: v for k, v in kwargs.items() if k in self.allowed_model_args474 }475 else:476 instance = {}477 instance[self.prompt_arg] = prompt478 instances.append(instance)479 480 predict_instances = [481 json_format.ParseDict(instance_dict, Value()) for instance_dict in instances482 ]483 return predict_instances484 485 def _generate(486 self,487 prompts: List[str],488 stop: Optional[List[str]] = None,489 run_manager: Optional[CallbackManagerForLLMRun] = None,490 **kwargs: Any,491 ) -> LLMResult:492 """Run the LLM on the given prompt and input."""493 instances = self._prepare_request(prompts, **kwargs)494 response = self.client.predict(endpoint=self.endpoint_path, instances=instances)495 return self._parse_response(response)496 497 def _parse_response(self, predictions: "Prediction") -> LLMResult:498 generations: List[List[Generation]] = []499 for result in predictions.predictions:500 generations.append(501 [502 Generation(text=self._parse_prediction(prediction))503 for prediction in result504 ]505 )506 return LLMResult(generations=generations)507 508 def _parse_prediction(self, prediction: Any) -> str:509 if isinstance(prediction, str):510 return prediction511 512 if self.result_arg:513 try:514 return prediction[self.result_arg]515 except KeyError:516 if isinstance(prediction, str):517 error_desc = (518 "Provided non-None `result_arg` (result_arg="519 f"{self.result_arg}). But got prediction of type "520 f"{type(prediction)} instead of dict. Most probably, you"521 "need to set `result_arg=None` during VertexAIModelGarden "522 "initialization."523 )524 raise ValueError(error_desc)525 else:526 raise ValueError(f"{self.result_arg} key not found in prediction!")527 528 return prediction529 530 async def _agenerate(531 self,532 prompts: List[str],533 stop: Optional[List[str]] = None,534 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,535 **kwargs: Any,536 ) -> LLMResult:537 """Run the LLM on the given prompt and input."""538 instances = self._prepare_request(prompts, **kwargs)539 response = await self.async_client.predict(540 endpoint=self.endpoint_path, instances=instances541 )542 return self._parse_response(response)543 