codekingpro/portable-devtools
114k
1import json2import logging3from typing import Any, Dict, List, Mapping, Optional, Set4 5import requests6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models.llms import LLM8from pydantic import Field9 10from langchain_community.llms.utils import enforce_stop_tokens11 12logger = logging.getLogger(__name__)13 14 15class Yuan2(LLM):16 """Yuan2.0 language models.17 18 Example:19 .. code-block:: python20 21 yuan_llm = Yuan2(22 infer_api="http://127.0.0.1:8000/yuan",23 max_tokens=1024,24 temp=1.0,25 top_p=0.9,26 top_k=40,27 )28 print(yuan_llm)29 print(yuan_llm.invoke("你是谁?"))30 """31 32 infer_api: str = "http://127.0.0.1:8000/yuan"33 """Yuan2.0 inference api"""34 35 max_tokens: int = Field(1024, alias="max_token")36 """Token context window."""37 38 temp: Optional[float] = 0.739 """The temperature to use for sampling."""40 41 top_p: Optional[float] = 0.942 """The top-p value to use for sampling."""43 44 top_k: Optional[int] = 045 """The top-k value to use for sampling."""46 47 do_sample: bool = False48 """The do_sample is a Boolean value that determines whether 49 to use the sampling method during text generation.50 """51 52 echo: Optional[bool] = False53 """Whether to echo the prompt."""54 55 stop: Optional[List[str]] = []56 """A list of strings to stop generation when encountered."""57 58 repeat_last_n: Optional[int] = 6459 "Last n tokens to penalize"60 61 repeat_penalty: Optional[float] = 1.1862 """The penalty to apply to repeated tokens."""63 64 streaming: bool = False65 """Whether to stream the results or not."""66 67 history: List[str] = []68 """History of the conversation"""69 70 use_history: bool = False71 """Whether to use history or not"""72 73 def __init__(self, **kwargs: Any) -> None:74 """Initialize the Yuan2 class."""75 super().__init__(**kwargs)76 77 if (self.top_p or 0) > 0 and (self.top_k or 0) > 0:78 logger.warning(79 "top_p and top_k cannot be set simultaneously. "80 "set top_k to 0 instead..."81 )82 self.top_k = 083 84 @property85 def _llm_type(self) -> str:86 return "Yuan2.0"87 88 @staticmethod89 def _model_param_names() -> Set[str]:90 return {91 "max_tokens",92 "temp",93 "top_k",94 "top_p",95 "do_sample",96 }97 98 def _default_params(self) -> Dict[str, Any]:99 return {100 "do_sample": self.do_sample,101 "infer_api": self.infer_api,102 "max_tokens": self.max_tokens,103 "repeat_penalty": self.repeat_penalty,104 "temp": self.temp,105 "top_k": self.top_k,106 "top_p": self.top_p,107 "use_history": self.use_history,108 }109 110 @property111 def _identifying_params(self) -> Mapping[str, Any]:112 """Get the identifying parameters."""113 return {114 "model": self._llm_type,115 **self._default_params(),116 **{117 k: v for k, v in self.__dict__.items() if k in self._model_param_names()118 },119 }120 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 a Yuan2.0 LLM inference endpoint.129 130 Args:131 prompt: The prompt to pass into the model.132 stop: Optional list of stop words to use when generating.133 134 Returns:135 The string generated by the model.136 137 Example:138 .. code-block:: python139 140 response = yuan_llm.invoke("你能做什么?")141 """142 143 if self.use_history:144 self.history.append(prompt)145 input = "<n>".join(self.history)146 else:147 input = prompt148 149 headers = {"Content-Type": "application/json"}150 151 data = json.dumps(152 {153 "ques_list": [{"id": "000", "ques": input}],154 "tokens_to_generate": self.max_tokens,155 "temperature": self.temp,156 "top_p": self.top_p,157 "top_k": self.top_k,158 "do_sample": self.do_sample,159 }160 )161 162 logger.debug("Yuan2.0 prompt:", input)163 164 # call api165 try:166 response = requests.put(self.infer_api, headers=headers, data=data)167 except requests.exceptions.RequestException as e:168 raise ValueError(f"Error raised by inference api: {e}")169 170 logger.debug(f"Yuan2.0 response: {response}")171 172 if response.status_code != 200:173 raise ValueError(f"Failed with response: {response}")174 try:175 resp = response.json()176 177 if resp["errCode"] != "0":178 raise ValueError(179 f"Failed with error code [{resp['errCode']}], "180 f"error message: [{resp['exceptionMsg']}]"181 )182 183 if "resData" in resp:184 if len(resp["resData"]["output"]) >= 0:185 generate_text = resp["resData"]["output"][0]["ans"]186 else:187 raise ValueError("No output found in response.")188 else:189 raise ValueError("No resData found in response.")190 191 except requests.exceptions.JSONDecodeError as e:192 raise ValueError(193 f"Error raised during decoding response from inference api: {e}."194 f"\nResponse: {response.text}"195 )196 197 if stop is not None:198 generate_text = enforce_stop_tokens(generate_text, stop)199 200 # support multi-turn chat201 if self.use_history:202 self.history.append(generate_text)203 204 logger.debug(f"history: {self.history}")205 return generate_text206 