Team Ai
Datasetpublic

codekingpro/portable-devtools

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