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