Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
yuan2.py206 linesDownload Raw Back to llms
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 
codekingpro/portable-devtools · Team Ai