codekingpro/portable-devtools
114k
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 