codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from typing import Any, Dict, Iterator, List, Optional4 5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.outputs import GenerationChunk8from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init9from pydantic import BaseModel, ConfigDict, Field, SecretStr10 11 12class VolcEngineMaasBase(BaseModel):13 """Base class for VolcEngineMaas models."""14 15 model_config = ConfigDict(protected_namespaces=())16 17 client: Any = None18 19 volc_engine_maas_ak: Optional[SecretStr] = None20 """access key for volc engine"""21 volc_engine_maas_sk: Optional[SecretStr] = None22 """secret key for volc engine"""23 24 endpoint: Optional[str] = "maas-api.ml-platform-cn-beijing.volces.com"25 """Endpoint of the VolcEngineMaas LLM."""26 27 region: Optional[str] = "Region"28 """Region of the VolcEngineMaas LLM."""29 30 model: str = "skylark-lite-public"31 """Model name. you could check this model details here 32 https://www.volcengine.com/docs/82379/113318733 and you could choose other models by change this field"""34 model_version: Optional[str] = None35 """Model version. Only used in moonshot large language model. 36 you could check details here https://www.volcengine.com/docs/82379/1158281"""37 38 top_p: Optional[float] = 0.839 """Total probability mass of tokens to consider at each step."""40 41 temperature: Optional[float] = 0.9542 """A non-negative float that tunes the degree of randomness in generation."""43 44 model_kwargs: Dict[str, Any] = Field(default_factory=dict)45 """model special arguments, you could check detail on model page"""46 47 streaming: bool = False48 """Whether to stream the results."""49 50 connect_timeout: Optional[int] = 6051 """Timeout for connect to volc engine maas endpoint. Default is 60 seconds."""52 53 read_timeout: Optional[int] = 6054 """Timeout for read response from volc engine maas endpoint. 55 Default is 60 seconds."""56 57 @pre_init58 def validate_environment(cls, values: Dict) -> Dict:59 volc_engine_maas_ak = convert_to_secret_str(60 get_from_dict_or_env(values, "volc_engine_maas_ak", "VOLC_ACCESSKEY")61 )62 volc_engine_maas_sk = convert_to_secret_str(63 get_from_dict_or_env(values, "volc_engine_maas_sk", "VOLC_SECRETKEY")64 )65 endpoint = values["endpoint"]66 if values["endpoint"] is not None and values["endpoint"] != "":67 endpoint = values["endpoint"]68 try:69 from volcengine.maas import MaasService70 71 maas = MaasService(72 endpoint,73 values["region"],74 connection_timeout=values["connect_timeout"],75 socket_timeout=values["read_timeout"],76 )77 maas.set_ak(volc_engine_maas_ak.get_secret_value())78 maas.set_sk(volc_engine_maas_sk.get_secret_value())79 80 values["volc_engine_maas_ak"] = volc_engine_maas_ak81 values["volc_engine_maas_sk"] = volc_engine_maas_sk82 values["client"] = maas83 except ImportError:84 raise ImportError(85 "volcengine package not found, please install it with "86 "`pip install volcengine`"87 )88 return values89 90 @property91 def _default_params(self) -> Dict[str, Any]:92 """Get the default parameters for calling VolcEngineMaas API."""93 normal_params = {94 "top_p": self.top_p,95 "temperature": self.temperature,96 }97 98 return {**normal_params, **self.model_kwargs}99 100 101class VolcEngineMaasLLM(LLM, VolcEngineMaasBase):102 """volc engine maas hosts a plethora of models.103 You can utilize these models through this class.104 105 To use, you should have the ``volcengine`` python package installed.106 and set access key and secret key by environment variable or direct pass those to107 this class.108 access key, secret key are required parameters which you could get help109 https://www.volcengine.com/docs/6291/65568110 111 In order to use them, it is necessary to install the 'volcengine' Python package.112 The access key and secret key must be set either via environment variables or113 passed directly to this class.114 access key and secret key are mandatory parameters for which assistance can be115 sought at https://www.volcengine.com/docs/6291/65568.116 117 Example:118 .. code-block:: python119 120 from langchain_community.llms import VolcEngineMaasLLM121 model = VolcEngineMaasLLM(model="skylark-lite-public",122 volc_engine_maas_ak="your_ak",123 volc_engine_maas_sk="your_sk")124 """125 126 @property127 def _llm_type(self) -> str:128 """Return type of llm."""129 return "volc-engine-maas-llm"130 131 def _convert_prompt_msg_params(132 self,133 prompt: str,134 **kwargs: Any,135 ) -> dict:136 model_req = {137 "model": {138 "name": self.model,139 }140 }141 if self.model_version is not None:142 model_req["model"]["version"] = self.model_version143 144 return {145 **model_req,146 "messages": [{"role": "user", "content": prompt}],147 "parameters": {**self._default_params, **kwargs},148 }149 150 def _call(151 self,152 prompt: str,153 stop: Optional[List[str]] = None,154 run_manager: Optional[CallbackManagerForLLMRun] = None,155 **kwargs: Any,156 ) -> str:157 if self.streaming:158 completion = ""159 for chunk in self._stream(prompt, stop, run_manager, **kwargs):160 completion += chunk.text161 return completion162 params = self._convert_prompt_msg_params(prompt, **kwargs)163 response = self.client.chat(params)164 165 return response.get("choice", {}).get("message", {}).get("content", "")166 167 def _stream(168 self,169 prompt: str,170 stop: Optional[List[str]] = None,171 run_manager: Optional[CallbackManagerForLLMRun] = None,172 **kwargs: Any,173 ) -> Iterator[GenerationChunk]:174 params = self._convert_prompt_msg_params(prompt, **kwargs)175 for res in self.client.stream_chat(params):176 if res:177 chunk = GenerationChunk(178 text=res.get("choice", {}).get("message", {}).get("content", "")179 )180 if run_manager:181 run_manager.on_llm_new_token(chunk.text, chunk=chunk)182 yield chunk183 