codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import json4from typing import (5 TYPE_CHECKING,6 Any,7 AsyncIterator,8 Dict,9 Generator,10 Iterator,11 List,12 Mapping,13 Optional,14 Union,15)16 17import aiohttp18import requests19from langchain_core.callbacks import (20 AsyncCallbackManagerForLLMRun,21 CallbackManagerForLLMRun,22)23from langchain_core.language_models.llms import LLM24from langchain_core.outputs import GenerationChunk25 26if TYPE_CHECKING:27 from xinference.client import RESTfulChatModelHandle, RESTfulGenerateModelHandle28 from xinference.model.llm.core import LlamaCppGenerateConfig29 30 31class Xinference(LLM):32 """`Xinference` large-scale model inference service.33 34 To use, you should have the xinference library installed:35 36 .. code-block:: bash37 38 pip install "xinference[all]"39 40 If you're simply using the services provided by Xinference, you can utilize the xinference_client package:41 42 .. code-block:: bash43 44 pip install xinference_client45 46 Check out: https://github.com/xorbitsai/inference47 To run, you need to start a Xinference supervisor on one server and Xinference workers on the other servers48 49 Example:50 To start a local instance of Xinference, run51 52 .. code-block:: bash53 54 $ xinference55 56 You can also deploy Xinference in a distributed cluster. Here are the steps:57 58 Starting the supervisor:59 60 .. code-block:: bash61 62 $ xinference-supervisor63 64 Starting the worker:65 66 .. code-block:: bash67 68 $ xinference-worker69 70 Then, launch a model using command line interface (CLI).71 72 Example:73 74 .. code-block:: bash75 76 $ xinference launch -n orca -s 3 -q q4_077 78 It will return a model UID. Then, you can use Xinference with LangChain.79 80 Example:81 82 .. code-block:: python83 84 from langchain_community.llms import Xinference85 86 llm = Xinference(87 server_url="http://0.0.0.0:9997",88 model_uid = {model_uid} # replace model_uid with the model UID return from launching the model89 )90 91 llm.invoke(92 prompt="Q: where can we visit in the capital of France? A:",93 generate_config={"max_tokens": 1024, "stream": True},94 )95 96 Example:97 98 .. code-block:: python99 100 from langchain_community.llms import Xinference101 from langchain_classic.prompts import PromptTemplate102 103 llm = Xinference(104 server_url="http://0.0.0.0:9997",105 model_uid={model_uid}, # replace model_uid with the model UID return from launching the model106 stream=True107 )108 prompt = PromptTemplate(109 input=['country'],110 template="Q: where can we visit in the capital of {country}? A:"111 )112 chain = prompt | llm113 chain.stream(input={'country': 'France'})114 115 116 To view all the supported builtin models, run:117 118 .. code-block:: bash119 120 $ xinference list --all121 122 """ # noqa: E501123 124 client: Optional[Any] = None125 server_url: Optional[str]126 """URL of the xinference server"""127 model_uid: Optional[str]128 """UID of the launched model"""129 model_kwargs: Dict[str, Any]130 """Keyword arguments to be passed to xinference.LLM"""131 132 def __init__(133 self,134 server_url: Optional[str] = None,135 model_uid: Optional[str] = None,136 api_key: Optional[str] = None,137 **model_kwargs: Any,138 ):139 try:140 from xinference.client import RESTfulClient141 except ImportError:142 try:143 from xinference_client import RESTfulClient144 except ImportError as e:145 raise ImportError(146 "Could not import RESTfulClient from xinference. Please install it"147 " with `pip install xinference` or `pip install xinference_client`."148 ) from e149 150 model_kwargs = model_kwargs or {}151 152 super().__init__(153 **{ # type: ignore[arg-type]154 "server_url": server_url,155 "model_uid": model_uid,156 "model_kwargs": model_kwargs,157 }158 )159 160 if self.server_url is None:161 raise ValueError("Please provide server URL")162 163 if self.model_uid is None:164 raise ValueError("Please provide the model UID")165 166 self._headers: Dict[str, str] = {}167 self._cluster_authed = False168 self._check_cluster_authenticated()169 if api_key is not None and self._cluster_authed:170 self._headers["Authorization"] = f"Bearer {api_key}"171 172 self.client = RESTfulClient(server_url, api_key)173 174 @property175 def _llm_type(self) -> str:176 """Return type of llm."""177 return "xinference"178 179 @property180 def _identifying_params(self) -> Mapping[str, Any]:181 """Get the identifying parameters."""182 return {183 **{"server_url": self.server_url},184 **{"model_uid": self.model_uid},185 **{"model_kwargs": self.model_kwargs},186 }187 188 def _check_cluster_authenticated(self) -> None:189 url = f"{self.server_url}/v1/cluster/auth"190 response = requests.get(url)191 if response.status_code == 404:192 self._cluster_authed = False193 else:194 if response.status_code != 200:195 raise RuntimeError(196 f"Failed to get cluster information, "197 f"detail: {response.json()['detail']}"198 )199 response_data = response.json()200 self._cluster_authed = bool(response_data["auth"])201 202 def _call(203 self,204 prompt: str,205 stop: Optional[List[str]] = None,206 run_manager: Optional[CallbackManagerForLLMRun] = None,207 **kwargs: Any,208 ) -> str:209 """Call the xinference model and return the output.210 211 Args:212 prompt: The prompt to use for generation.213 stop: Optional list of stop words to use when generating.214 generate_config: Optional dictionary for the configuration used for215 generation.216 217 Returns:218 The generated string by the model.219 """220 if self.client is None:221 raise ValueError("Client is not initialized!")222 model = self.client.get_model(self.model_uid)223 224 generate_config: "LlamaCppGenerateConfig" = kwargs.get("generate_config", {})225 226 generate_config = {**self.model_kwargs, **generate_config}227 228 if stop:229 generate_config["stop"] = stop230 231 if generate_config and generate_config.get("stream"):232 combined_text_output = ""233 for token in self._stream_generate(234 model=model,235 prompt=prompt,236 run_manager=run_manager,237 generate_config=generate_config,238 ):239 combined_text_output += token240 return combined_text_output241 242 else:243 completion = model.generate(prompt=prompt, generate_config=generate_config)244 return completion["choices"][0]["text"]245 246 def _stream_generate(247 self,248 model: Union["RESTfulGenerateModelHandle", "RESTfulChatModelHandle"],249 prompt: str,250 run_manager: Optional[CallbackManagerForLLMRun] = None,251 generate_config: Optional["LlamaCppGenerateConfig"] = None,252 ) -> Generator[str, None, None]:253 """254 Args:255 prompt: The prompt to use for generation.256 model: The model used for generation.257 stop: Optional list of stop words to use when generating.258 generate_config: Optional dictionary for the configuration used for259 generation.260 261 Yields:262 A string token.263 """264 streaming_response = model.generate(265 prompt=prompt, generate_config=generate_config266 )267 for chunk in streaming_response:268 if isinstance(chunk, dict):269 choices = chunk.get("choices", [])270 if choices:271 choice = choices[0]272 if isinstance(choice, dict):273 token = choice.get("text", "")274 log_probs = choice.get("logprobs")275 if run_manager:276 run_manager.on_llm_new_token(277 token=token, verbose=self.verbose, log_probs=log_probs278 )279 yield token280 281 def _stream(282 self,283 prompt: str,284 stop: Optional[List[str]] = None,285 run_manager: Optional[CallbackManagerForLLMRun] = None,286 **kwargs: Any,287 ) -> Iterator[GenerationChunk]:288 generate_config = kwargs.get("generate_config", {})289 generate_config = {**self.model_kwargs, **generate_config}290 if stop:291 generate_config["stop"] = stop292 for stream_resp in self._create_generate_stream(prompt, generate_config):293 if stream_resp:294 chunk = self._stream_response_to_generation_chunk(stream_resp)295 if run_manager:296 run_manager.on_llm_new_token(297 chunk.text,298 verbose=self.verbose,299 )300 yield chunk301 302 def _create_generate_stream(303 self, prompt: str, generate_config: Optional[Dict[str, List[str]]] = None304 ) -> Iterator[str]:305 if self.client is None:306 raise ValueError("Client is not initialized!")307 model = self.client.get_model(self.model_uid)308 yield from model.generate(prompt=prompt, generate_config=generate_config)309 310 @staticmethod311 def _stream_response_to_generation_chunk(312 stream_response: str,313 ) -> GenerationChunk:314 """Convert a stream response to a generation chunk."""315 token = ""316 if isinstance(stream_response, dict):317 choices = stream_response.get("choices", [])318 if choices:319 choice = choices[0]320 if isinstance(choice, dict):321 token = choice.get("text", "")322 323 return GenerationChunk(324 text=token,325 generation_info=dict(326 finish_reason=choice.get("finish_reason", None),327 logprobs=choice.get("logprobs", None),328 ),329 )330 else:331 raise TypeError("choice type error!")332 else:333 return GenerationChunk(text=token)334 else:335 raise TypeError("stream_response type error!")336 337 async def _astream(338 self,339 prompt: str,340 stop: Optional[List[str]] = None,341 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,342 **kwargs: Any,343 ) -> AsyncIterator[GenerationChunk]:344 generate_config = kwargs.get("generate_config", {})345 generate_config = {**self.model_kwargs, **generate_config}346 if stop:347 generate_config["stop"] = stop348 async for stream_resp in self._acreate_generate_stream(prompt, generate_config):349 if stream_resp:350 chunk = self._stream_response_to_generation_chunk(stream_resp)351 if run_manager:352 await run_manager.on_llm_new_token(353 chunk.text,354 verbose=self.verbose,355 )356 yield chunk357 358 async def _acreate_generate_stream(359 self, prompt: str, generate_config: Optional[Dict[str, List[str]]] = None360 ) -> AsyncIterator[str]:361 request_body: Dict[str, Any] = {"model": self.model_uid, "prompt": prompt}362 if generate_config is not None:363 for key, value in generate_config.items():364 request_body[key] = value365 366 stream = bool(generate_config and generate_config.get("stream"))367 async with aiohttp.ClientSession() as session:368 async with session.post(369 url=f"{self.server_url}/v1/completions",370 json=request_body,371 ) as response:372 if response.status != 200:373 if response.status == 404:374 raise FileNotFoundError(375 "astream call failed with status code 404."376 )377 else:378 optional_detail = response.text379 raise ValueError(380 f"astream call failed with status code {response.status}."381 f" Details: {optional_detail}"382 )383 384 async for line in response.content:385 if not stream:386 yield json.loads(line)387 else:388 json_str = line.decode("utf-8")389 if line.startswith(b"data:"):390 json_str = json_str[len(b"data:") :].strip()391 if not json_str:392 continue393 yield json.loads(json_str)394 