codekingpro/portable-devtools
114k
1import logging2from typing import Any, Dict, List, Optional, Union3 4from langchain_core._api.deprecation import deprecated5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.utils import get_from_dict_or_env8from pydantic import BaseModel, ConfigDict, model_validator9 10from langchain_community.llms.utils import enforce_stop_tokens11 12logger = logging.getLogger(__name__)13 14 15@deprecated(16 since="0.3.28",17 removal="1.0",18 alternative_import="langchain_predictionguard.PredictionGuard",19)20class PredictionGuard(LLM):21 """Prediction Guard large language models.22 23 To use, you should have the ``predictionguard`` python package installed, and the24 environment variable ``PREDICTIONGUARD_API_KEY`` set with your API key, or pass25 it as a named parameter to the constructor.26 27 Example:28 .. code-block:: python29 30 llm = PredictionGuard(31 model="Hermes-3-Llama-3.1-8B",32 predictionguard_api_key="your Prediction Guard API key",33 )34 """35 36 client: Any = None #: :meta private:37 38 model: Optional[str] = "Hermes-3-Llama-3.1-8B"39 """Model name to use."""40 41 max_tokens: Optional[int] = 25642 """Denotes the number of tokens to predict per generation."""43 44 temperature: Optional[float] = 0.7545 """A non-negative float that tunes the degree of randomness in generation."""46 47 top_p: Optional[float] = 0.148 """A non-negative float that controls the diversity of the generated tokens."""49 50 top_k: Optional[int] = None51 """The diversity of the generated text based on top-k sampling."""52 53 stop: Optional[List[str]] = None54 55 predictionguard_input: Optional[Dict[str, Union[str, bool]]] = None56 """The input check to run over the prompt before sending to the LLM."""57 58 predictionguard_output: Optional[Dict[str, bool]] = None59 """The output check to run the LLM output against."""60 61 predictionguard_api_key: Optional[str] = None62 """Prediction Guard API key."""63 64 model_config = ConfigDict(extra="forbid")65 66 @model_validator(mode="before")67 def validate_environment(cls, values: Dict) -> Dict:68 """Validate that the api_key and python package exists in environment."""69 pg_api_key = get_from_dict_or_env(70 values, "predictionguard_api_key", "PREDICTIONGUARD_API_KEY"71 )72 73 try:74 from predictionguard import PredictionGuard75 76 values["client"] = PredictionGuard(77 api_key=pg_api_key,78 )79 80 except ImportError:81 raise ImportError(82 "Could not import predictionguard python package. "83 "Please install it with `pip install predictionguard`."84 )85 86 return values87 88 @property89 def _identifying_params(self) -> Dict[str, Any]:90 """Get the identifying parameters."""91 return {"model": self.model}92 93 @property94 def _llm_type(self) -> str:95 """Return type of llm."""96 return "predictionguard"97 98 def _get_parameters(self, **kwargs: Any) -> Dict[str, Any]:99 # input kwarg conflicts with LanguageModelInput on BaseChatModel100 input = kwargs.pop("predictionguard_input", self.predictionguard_input)101 output = kwargs.pop("predictionguard_output", self.predictionguard_output)102 103 params = {104 **{105 "max_tokens": self.max_tokens,106 "temperature": self.temperature,107 "top_p": self.top_p,108 "top_k": self.top_k,109 "input": (110 input.model_dump() if isinstance(input, BaseModel) else input111 ),112 "output": (113 output.model_dump() if isinstance(output, BaseModel) else output114 ),115 },116 **kwargs,117 }118 119 return params120 121 def _call(122 self,123 prompt: str,124 stop: Optional[List[str]] = None,125 run_manager: Optional[CallbackManagerForLLMRun] = None,126 **kwargs: Any,127 ) -> str:128 """Call out to Prediction Guard's model API.129 Args:130 prompt: The prompt to pass into the model.131 Returns:132 The string generated by the model.133 Example:134 .. code-block:: python135 response = llm.invoke("Tell me a joke.")136 """137 138 params = self._get_parameters(**kwargs)139 140 stops = None141 if self.stop is not None and stop is not None:142 raise ValueError("`stop` found in both the input and default params.")143 elif self.stop is not None:144 stops = self.stop145 else:146 stops = stop147 148 response = self.client.completions.create(149 model=self.model,150 prompt=prompt,151 **params,152 )153 154 for res in response["choices"]:155 if res.get("status", "").startswith("error: "):156 err_msg = res["status"].removeprefix("error: ")157 raise ValueError(f"Error from PredictionGuard API: {err_msg}")158 159 text = response["choices"][0]["text"]160 161 # If stop tokens are provided, Prediction Guard's endpoint returns them.162 # In order to make this consistent with other endpoints, we strip them.163 if stops:164 text = enforce_stop_tokens(text, stops)165 166 return text167 