Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
chatglm3.py152 linesDownload Raw Back to llms
1import json2import logging3from typing import Any, List, Optional, Union4 5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.messages import (8    AIMessage,9    BaseMessage,10    FunctionMessage,11    HumanMessage,12    SystemMessage,13)14from pydantic import Field15 16from langchain_community.llms.utils import enforce_stop_tokens17 18logger = logging.getLogger(__name__)19HEADERS = {"Content-Type": "application/json"}20DEFAULT_TIMEOUT = 3021 22 23def _convert_message_to_dict(message: BaseMessage) -> dict:24    if isinstance(message, HumanMessage):25        message_dict = {"role": "user", "content": message.content}26    elif isinstance(message, AIMessage):27        message_dict = {"role": "assistant", "content": message.content}28    elif isinstance(message, SystemMessage):29        message_dict = {"role": "system", "content": message.content}30    elif isinstance(message, FunctionMessage):31        message_dict = {"role": "function", "content": message.content}32    else:33        raise ValueError(f"Got unknown type {message}")34    return message_dict35 36 37class ChatGLM3(LLM):38    """ChatGLM3 LLM service."""39 40    model_name: str = Field(default="chatglm3-6b", alias="model")41    endpoint_url: str = "http://127.0.0.1:8000/v1/chat/completions"42    """Endpoint URL to use."""43    model_kwargs: Optional[dict] = None44    """Keyword arguments to pass to the model."""45    max_tokens: int = 2000046    """Max token allowed to pass to the model."""47    temperature: float = 0.148    """LLM model temperature from 0 to 10."""49    top_p: float = 0.750    """Top P for nucleus sampling from 0 to 1"""51    prefix_messages: List[BaseMessage] = Field(default_factory=list)52    """Series of messages for Chat input."""53    streaming: bool = False54    """Whether to stream the results or not."""55    http_client: Union[Any, None] = None56    timeout: int = DEFAULT_TIMEOUT57 58    @property59    def _llm_type(self) -> str:60        return "chat_glm_3"61 62    @property63    def _invocation_params(self) -> dict:64        """Get the parameters used to invoke the model."""65        params = {66            "model": self.model_name,67            "temperature": self.temperature,68            "max_tokens": self.max_tokens,69            "top_p": self.top_p,70            "stream": self.streaming,71        }72        return {**params, **(self.model_kwargs or {})}73 74    @property75    def client(self) -> Any:76        import httpx77 78        return self.http_client or httpx.Client(timeout=self.timeout)79 80    def _get_payload(self, prompt: str) -> dict:81        params = self._invocation_params82        messages = self.prefix_messages + [HumanMessage(content=prompt)]83        params.update(84            {85                "messages": [_convert_message_to_dict(m) for m in messages],86            }87        )88        return params89 90    def _call(91        self,92        prompt: str,93        stop: Optional[List[str]] = None,94        run_manager: Optional[CallbackManagerForLLMRun] = None,95        **kwargs: Any,96    ) -> str:97        """Call out to a ChatGLM3 LLM inference endpoint.98 99        Args:100            prompt: The prompt to pass into the model.101            stop: Optional list of stop words to use when generating.102 103        Returns:104            The string generated by the model.105 106        Example:107            .. code-block:: python108 109                response = chatglm_llm.invoke("Who are you?")110        """111        import httpx112 113        payload = self._get_payload(prompt)114        logger.debug(f"ChatGLM3 payload: {payload}")115 116        try:117            response = self.client.post(118                self.endpoint_url, headers=HEADERS, json=payload119            )120        except httpx.NetworkError as e:121            raise ValueError(f"Error raised by inference endpoint: {e}")122 123        logger.debug(f"ChatGLM3 response: {response}")124 125        if response.status_code != 200:126            raise ValueError(f"Failed with response: {response}")127 128        try:129            parsed_response = response.json()130 131            if isinstance(parsed_response, dict):132                content_keys = "choices"133                if content_keys in parsed_response:134                    choices = parsed_response[content_keys]135                    if len(choices):136                        text = choices[0]["message"]["content"]137                else:138                    raise ValueError(f"No content in response : {parsed_response}")139            else:140                raise ValueError(f"Unexpected response type: {parsed_response}")141 142        except json.JSONDecodeError as e:143            raise ValueError(144                f"Error raised during decoding response from inference endpoint: {e}."145                f"\nResponse: {response.text}"146            )147 148        if stop is not None:149            text = enforce_stop_tokens(text, stop)150 151        return text152