codekingpro/portable-devtools
114k
1import logging2from typing import Any, Dict, List, Mapping, Optional, cast3 4import requests5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init8from pydantic import ConfigDict, Field, SecretStr, model_validator9 10from langchain_community.llms.utils import enforce_stop_tokens11 12logger = logging.getLogger(__name__)13 14 15class CerebriumAI(LLM):16 """CerebriumAI large language models.17 18 To use, you should have the ``cerebrium`` python package installed.19 You should also have the environment variable ``CEREBRIUMAI_API_KEY``20 set with your API key or pass it as a named argument in the constructor.21 22 Any parameters that are valid to be passed to the call can be passed23 in, even if not explicitly saved on this class.24 25 Example:26 .. code-block:: python27 28 from langchain_community.llms import CerebriumAI29 cerebrium = CerebriumAI(endpoint_url="", cerebriumai_api_key="my-api-key")30 31 """32 33 endpoint_url: str = ""34 """model endpoint to use"""35 36 model_kwargs: Dict[str, Any] = Field(default_factory=dict)37 """Holds any model parameters valid for `create` call not38 explicitly specified."""39 40 cerebriumai_api_key: Optional[SecretStr] = None41 42 model_config = ConfigDict(43 extra="forbid",44 )45 46 @model_validator(mode="before")47 @classmethod48 def build_extra(cls, values: Dict[str, Any]) -> Any:49 """Build extra kwargs from additional params that were passed in."""50 all_required_field_names = set(list(cls.model_fields.keys()))51 52 extra = values.get("model_kwargs", {})53 for field_name in list(values):54 if field_name not in all_required_field_names:55 if field_name in extra:56 raise ValueError(f"Found {field_name} supplied twice.")57 logger.warning(58 f"""{field_name} was transferred to model_kwargs.59 Please confirm that {field_name} is what you intended."""60 )61 extra[field_name] = values.pop(field_name)62 values["model_kwargs"] = extra63 return values64 65 @pre_init66 def validate_environment(cls, values: Dict) -> Dict:67 """Validate that api key and python package exists in environment."""68 cerebriumai_api_key = convert_to_secret_str(69 get_from_dict_or_env(values, "cerebriumai_api_key", "CEREBRIUMAI_API_KEY")70 )71 values["cerebriumai_api_key"] = cerebriumai_api_key72 return values73 74 @property75 def _identifying_params(self) -> Mapping[str, Any]:76 """Get the identifying parameters."""77 return {78 **{"endpoint_url": self.endpoint_url},79 **{"model_kwargs": self.model_kwargs},80 }81 82 @property83 def _llm_type(self) -> str:84 """Return type of llm."""85 return "cerebriumai"86 87 def _call(88 self,89 prompt: str,90 stop: Optional[List[str]] = None,91 run_manager: Optional[CallbackManagerForLLMRun] = None,92 **kwargs: Any,93 ) -> str:94 headers: Dict = {95 "Authorization": cast(96 SecretStr, self.cerebriumai_api_key97 ).get_secret_value(),98 "Content-Type": "application/json",99 }100 params = self.model_kwargs or {}101 payload = {"prompt": prompt, **params, **kwargs}102 response = requests.post(self.endpoint_url, json=payload, headers=headers)103 if response.status_code == 200:104 data = response.json()105 text = data["result"]106 if stop is not None:107 # I believe this is required since the stop tokens108 # are not enforced by the model parameters109 text = enforce_stop_tokens(text, stop)110 return text111 else:112 response.raise_for_status()113 return ""114 