Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
google_palm.py246 linesDownload Raw Back to llms
1from __future__ import annotations2 3from typing import Any, Dict, Iterator, List, Optional4 5from langchain_core._api.deprecation import deprecated6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models import LanguageModelInput8from langchain_core.outputs import Generation, GenerationChunk, LLMResult9from langchain_core.utils import get_from_dict_or_env, pre_init10from pydantic import BaseModel, SecretStr11 12from langchain_community.llms import BaseLLM13from langchain_community.utilities.vertexai import create_retry_decorator14 15 16def completion_with_retry(17    llm: GooglePalm,18    prompt: LanguageModelInput,19    is_gemini: bool = False,20    stream: bool = False,21    run_manager: Optional[CallbackManagerForLLMRun] = None,22    **kwargs: Any,23) -> Any:24    """Use tenacity to retry the completion call."""25    retry_decorator = create_retry_decorator(26        llm, max_retries=llm.max_retries, run_manager=run_manager27    )28 29    @retry_decorator30    def _completion_with_retry(31        prompt: LanguageModelInput, is_gemini: bool, stream: bool, **kwargs: Any32    ) -> Any:33        generation_config = kwargs.get("generation_config", {})34        if is_gemini:35            return llm.client.generate_content(36                contents=prompt, stream=stream, generation_config=generation_config37            )38        return llm.client.generate_text(prompt=prompt, **kwargs)39 40    return _completion_with_retry(41        prompt=prompt, is_gemini=is_gemini, stream=stream, **kwargs42    )43 44 45def _is_gemini_model(model_name: str) -> bool:46    return "gemini" in model_name47 48 49def _strip_erroneous_leading_spaces(text: str) -> str:50    """Strip erroneous leading spaces from text.51 52    The PaLM API will sometimes erroneously return a single leading space in all53    lines > 1. This function strips that space.54    """55    has_leading_space = all(not line or line[0] == " " for line in text.split("\n")[1:])56    if has_leading_space:57        return text.replace("\n ", "\n")58    else:59        return text60 61 62@deprecated("0.0.12", alternative_import="langchain_google_genai.GoogleGenerativeAI")63class GooglePalm(BaseLLM, BaseModel):64    """65    DEPRECATED: Use `langchain_google_genai.GoogleGenerativeAI` instead.66 67    Google PaLM models.68    """69 70    client: Any  #: :meta private:71    google_api_key: Optional[SecretStr]72    model_name: str = "models/text-bison-001"73    """Model name to use."""74    temperature: float = 0.775    """Run inference with this temperature. Must be in the closed interval76       [0.0, 1.0]."""77    top_p: Optional[float] = None78    """Decode using nucleus sampling: consider the smallest set of tokens whose79       probability sum is at least top_p. Must be in the closed interval [0.0, 1.0]."""80    top_k: Optional[int] = None81    """Decode using top-k sampling: consider the set of top_k most probable tokens.82       Must be positive."""83    max_output_tokens: Optional[int] = None84    """Maximum number of tokens to include in a candidate. Must be greater than zero.85       If unset, will default to 64."""86    n: int = 187    """Number of chat completions to generate for each prompt. Note that the API may88       not return the full n completions if duplicates are generated."""89    max_retries: int = 690    """The maximum number of retries to make when generating."""91 92    @property93    def is_gemini(self) -> bool:94        """Returns whether a model is belongs to a Gemini family or not."""95        return _is_gemini_model(self.model_name)96 97    @property98    def lc_secrets(self) -> Dict[str, str]:99        return {"google_api_key": "GOOGLE_API_KEY"}100 101    @classmethod102    def is_lc_serializable(self) -> bool:103        return True104 105    @classmethod106    def get_lc_namespace(cls) -> List[str]:107        """Get the namespace of the langchain object."""108        return ["langchain", "llms", "google_palm"]109 110    @pre_init111    def validate_environment(cls, values: Dict) -> Dict:112        """Validate api key, python package exists."""113        google_api_key = get_from_dict_or_env(114            values, "google_api_key", "GOOGLE_API_KEY"115        )116        model_name = values["model_name"]117        try:118            import google.generativeai as genai119 120            if isinstance(google_api_key, SecretStr):121                google_api_key = google_api_key.get_secret_value()122 123            genai.configure(api_key=google_api_key)124 125            if _is_gemini_model(model_name):126                values["client"] = genai.GenerativeModel(model_name=model_name)127            else:128                values["client"] = genai129        except ImportError:130            raise ImportError(131                "Could not import google-generativeai python package. "132                "Please install it with `pip install google-generativeai`."133            )134 135        if values["temperature"] is not None and not 0 <= values["temperature"] <= 1:136            raise ValueError("temperature must be in the range [0.0, 1.0]")137 138        if values["top_p"] is not None and not 0 <= values["top_p"] <= 1:139            raise ValueError("top_p must be in the range [0.0, 1.0]")140 141        if values["top_k"] is not None and values["top_k"] <= 0:142            raise ValueError("top_k must be positive")143 144        if values["max_output_tokens"] is not None and values["max_output_tokens"] <= 0:145            raise ValueError("max_output_tokens must be greater than zero")146 147        return values148 149    def _generate(150        self,151        prompts: List[str],152        stop: Optional[List[str]] = None,153        run_manager: Optional[CallbackManagerForLLMRun] = None,154        **kwargs: Any,155    ) -> LLMResult:156        generations: List[List[Generation]] = []157        generation_config = {158            "stop_sequences": stop,159            "temperature": self.temperature,160            "top_p": self.top_p,161            "top_k": self.top_k,162            "max_output_tokens": self.max_output_tokens,163            "candidate_count": self.n,164        }165        for prompt in prompts:166            if self.is_gemini:167                res = completion_with_retry(168                    self,169                    prompt=prompt,170                    stream=False,171                    is_gemini=True,172                    run_manager=run_manager,173                    generation_config=generation_config,174                )175                candidates = [176                    "".join([p.text for p in c.content.parts]) for c in res.candidates177                ]178                generations.append([Generation(text=c) for c in candidates])179            else:180                res = completion_with_retry(181                    self,182                    model=self.model_name,183                    prompt=prompt,184                    stream=False,185                    is_gemini=False,186                    run_manager=run_manager,187                    **generation_config,188                )189                prompt_generations = []190                for candidate in res.candidates:191                    raw_text = candidate["output"]192                    stripped_text = _strip_erroneous_leading_spaces(raw_text)193                    prompt_generations.append(Generation(text=stripped_text))194                generations.append(prompt_generations)195 196        return LLMResult(generations=generations)197 198    def _stream(199        self,200        prompt: str,201        stop: Optional[List[str]] = None,202        run_manager: Optional[CallbackManagerForLLMRun] = None,203        **kwargs: Any,204    ) -> Iterator[GenerationChunk]:205        generation_config = kwargs.get("generation_config", {})206        if stop:207            generation_config["stop_sequences"] = stop208        for stream_resp in completion_with_retry(209            self,210            prompt,211            stream=True,212            is_gemini=True,213            run_manager=run_manager,214            generation_config=generation_config,215            **kwargs,216        ):217            chunk = GenerationChunk(text=stream_resp.text)218            if run_manager:219                run_manager.on_llm_new_token(220                    stream_resp.text,221                    chunk=chunk,222                    verbose=self.verbose,223                )224            yield chunk225 226    @property227    def _llm_type(self) -> str:228        """Return type of llm."""229        return "google_palm"230 231    def get_num_tokens(self, text: str) -> int:232        """Get the number of tokens present in the text.233 234        Useful for checking if an input will fit in a model's context window.235 236        Args:237            text: The string input to tokenize.238 239        Returns:240            The integer number of tokens in the text.241        """242        if self.is_gemini:243            raise ValueError("Counting tokens is not yet supported!")244        result = self.client.count_text_tokens(model=self.model_name, prompt=text)245        return result["token_count"]246 
codekingpro/portable-devtools · Team Ai