codekingpro/portable-devtools
114k
1import json2from typing import Any, Dict, Iterator, List, Optional, Tuple, Union3 4import requests5from langchain_core.callbacks.manager import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.outputs import GenerationChunk8from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env9from pydantic import Field, SecretStr10from requests import Response11 12 13class SambaStudio(LLM):14 """15 SambaStudio large language models.16 17 Setup:18 To use, you should have the environment variables19 ``SAMBASTUDIO_URL`` set with your SambaStudio environment URL.20 ``SAMBASTUDIO_API_KEY`` set with your SambaStudio endpoint API key.21 https://sambanova.ai/products/enterprise-ai-platform-sambanova-suite22 read extra documentation in https://docs.sambanova.ai/sambastudio/latest/index.html23 Example:24 .. code-block:: python25 from langchain_community.llms.sambanova import SambaStudio26 SambaStudio(27 sambastudio_url="your-SambaStudio-environment-URL",28 sambastudio_api_key="your-SambaStudio-API-key,29 model_kwargs={30 "model" : model or expert name (set for Bundle endpoints),31 "max_tokens" : max number of tokens to generate,32 "temperature" : model temperature,33 "top_p" : model top p,34 "top_k" : model top k,35 "do_sample" : wether to do sample36 "process_prompt": wether to process prompt37 (set for Bundle generic v1 and v2 endpoints)38 },39 )40 Key init args — completion params:41 model: str42 The name of the model to use, e.g., Meta-Llama-3-70B-Instruct-409643 (set for Bundle endpoints).44 streaming: bool45 Whether to use streaming handler when using non streaming methods46 model_kwargs: dict47 Extra Key word arguments to pass to the model:48 max_tokens: int49 max tokens to generate50 temperature: float51 model temperature52 top_p: float53 model top p54 top_k: int55 model top k56 do_sample: bool57 wether to do sample58 process_prompt:59 wether to process prompt60 (set for Bundle generic v1 and v2 endpoints)61 Key init args — client params:62 sambastudio_url: str63 SambaStudio endpoint Url64 sambastudio_api_key: str65 SambaStudio endpoint api key66 67 Instantiate:68 .. code-block:: python69 70 from langchain_community.llms import SambaStudio71 72 llm = SambaStudio=(73 sambastudio_url = set with your SambaStudio deployed endpoint URL,74 sambastudio_api_key = set with your SambaStudio deployed endpoint Key,75 model_kwargs = {76 "model" : model or expert name (set for Bundle endpoints),77 "max_tokens" : max number of tokens to generate,78 "temperature" : model temperature,79 "top_p" : model top p,80 "top_k" : model top k,81 "do_sample" : wether to do sample82 "process_prompt" : wether to process prompt83 (set for Bundle generic v1 and v2 endpoints)84 }85 )86 87 Invoke:88 .. code-block:: python89 prompt = "tell me a joke"90 response = llm.invoke(prompt)91 92 Stream:93 .. code-block:: python94 95 for chunk in llm.stream(prompt):96 print(chunk, end="", flush=True)97 98 Async:99 .. code-block:: python100 101 response = llm.ainvoke(prompt)102 await response103 104 """105 106 sambastudio_url: str = Field(default="")107 """SambaStudio Url"""108 109 sambastudio_api_key: SecretStr = Field(default=SecretStr(""))110 """SambaStudio api key"""111 112 base_url: str = Field(default="", exclude=True)113 """SambaStudio non streaming URL"""114 115 streaming_url: str = Field(default="", exclude=True)116 """SambaStudio streaming URL"""117 118 streaming: bool = Field(default=False)119 """Whether to use streaming handler when using non streaming methods"""120 121 model_kwargs: Optional[Dict[str, Any]] = None122 """Key word arguments to pass to the model."""123 124 class Config:125 populate_by_name = True126 127 @classmethod128 def is_lc_serializable(cls) -> bool:129 """Return whether this model can be serialized by Langchain."""130 return True131 132 @property133 def lc_secrets(self) -> Dict[str, str]:134 return {135 "sambastudio_url": "sambastudio_url",136 "sambastudio_api_key": "sambastudio_api_key",137 }138 139 @property140 def _identifying_params(self) -> Dict[str, Any]:141 """Return a dictionary of identifying parameters.142 143 This information is used by the LangChain callback system, which144 is used for tracing purposes make it possible to monitor LLMs.145 """146 return {"streaming": self.streaming, **{"model_kwargs": self.model_kwargs}}147 148 @property149 def _llm_type(self) -> str:150 """Return type of llm."""151 return "sambastudio-llm"152 153 def __init__(self, **kwargs: Any) -> None:154 """init and validate environment variables"""155 kwargs["sambastudio_url"] = get_from_dict_or_env(156 kwargs, "sambastudio_url", "SAMBASTUDIO_URL"157 )158 159 kwargs["sambastudio_api_key"] = convert_to_secret_str(160 get_from_dict_or_env(kwargs, "sambastudio_api_key", "SAMBASTUDIO_API_KEY")161 )162 kwargs["base_url"], kwargs["streaming_url"] = self._get_sambastudio_urls(163 kwargs["sambastudio_url"]164 )165 super().__init__(**kwargs)166 167 def _get_sambastudio_urls(self, url: str) -> Tuple[str, str]:168 """169 Get streaming and non streaming URLs from the given URL170 171 Args:172 url: string with sambastudio base or streaming endpoint url173 174 Returns:175 base_url: string with url to do non streaming calls176 streaming_url: string with url to do streaming calls177 """178 if "chat/completions" in url:179 base_url = url180 stream_url = url181 else:182 if "stream" in url:183 base_url = url.replace("stream/", "")184 stream_url = url185 else:186 base_url = url187 if "generic" in url:188 stream_url = "generic/stream".join(url.split("generic"))189 else:190 raise ValueError("Unsupported URL")191 return base_url, stream_url192 193 def _get_tuning_params(self, stop: Optional[List[str]] = None) -> Dict[str, Any]:194 """195 Get the tuning parameters to use when calling the LLM.196 197 Args:198 stop: Stop words to use when generating. Model output is cut off at the199 first occurrence of any of the stop substrings.200 201 Returns:202 The tuning parameters in the format required by api to use203 """204 if stop is None:205 stop = []206 207 # get the parameters to use when calling the LLM.208 _model_kwargs = self.model_kwargs or {}209 210 # handle the case where stop sequences are send in the invocation211 # and stop sequences has been also set in the model parameters212 _stop_sequences = _model_kwargs.get("stop_sequences", []) + stop213 if len(_stop_sequences) > 0:214 _model_kwargs["stop_sequences"] = _stop_sequences215 216 # set the parameters structure depending of the API217 if "chat/completions" in self.sambastudio_url:218 if "select_expert" in _model_kwargs.keys():219 _model_kwargs["model"] = _model_kwargs.pop("select_expert")220 if "max_tokens_to_generate" in _model_kwargs.keys():221 _model_kwargs["max_tokens"] = _model_kwargs.pop(222 "max_tokens_to_generate"223 )224 if "process_prompt" in _model_kwargs.keys():225 _model_kwargs.pop("process_prompt")226 tuning_params = _model_kwargs227 228 elif "api/v2/predict/generic" in self.sambastudio_url:229 if "model" in _model_kwargs.keys():230 _model_kwargs["select_expert"] = _model_kwargs.pop("model")231 if "max_tokens" in _model_kwargs.keys():232 _model_kwargs["max_tokens_to_generate"] = _model_kwargs.pop(233 "max_tokens"234 )235 tuning_params = _model_kwargs236 237 elif "api/predict/generic" in self.sambastudio_url:238 if "model" in _model_kwargs.keys():239 _model_kwargs["select_expert"] = _model_kwargs.pop("model")240 if "max_tokens" in _model_kwargs.keys():241 _model_kwargs["max_tokens_to_generate"] = _model_kwargs.pop(242 "max_tokens"243 )244 245 tuning_params = {246 k: {"type": type(v).__name__, "value": str(v)}247 for k, v in (_model_kwargs.items())248 }249 250 else:251 raise ValueError(252 f"Unsupported URL{self.sambastudio_url}"253 "only openai, generic v1 and generic v2 APIs are supported"254 )255 256 return tuning_params257 258 def _handle_request(259 self,260 prompt: Union[List[str], str],261 stop: Optional[List[str]] = None,262 streaming: Optional[bool] = False,263 ) -> Response:264 """265 Performs a post request to the LLM API.266 267 Args:268 prompt: The prompt to pass into the model269 stop: list of stop tokens270 streaming: wether to do a streaming call271 272 Returns:273 A request Response object274 """275 276 if isinstance(prompt, str):277 prompt = [prompt]278 279 params = self._get_tuning_params(stop)280 281 # create request payload for openAI v1 API282 if "chat/completions" in self.sambastudio_url:283 messages_dict = [{"role": "user", "content": prompt[0]}]284 data = {"messages": messages_dict, "stream": streaming, **params}285 data = {key: value for key, value in data.items() if value is not None}286 headers = {287 "Authorization": f"Bearer "288 f"{self.sambastudio_api_key.get_secret_value()}",289 "Content-Type": "application/json",290 }291 292 # create request payload for generic v1 API293 elif "api/v2/predict/generic" in self.sambastudio_url:294 if params.get("process_prompt", False):295 prompt = json.dumps(296 {297 "conversation_id": "sambaverse-conversation-id",298 "messages": [299 {"message_id": None, "role": "user", "content": prompt[0]}300 ],301 }302 )303 else:304 prompt = prompt[0]305 items = [{"id": "item0", "value": prompt}]306 params = {key: value for key, value in params.items() if value is not None}307 data = {"items": items, "params": params}308 headers = {"key": self.sambastudio_api_key.get_secret_value()}309 310 # create request payload for generic v1 API311 elif "api/predict/generic" in self.sambastudio_url:312 if params.get("process_prompt", False):313 if params["process_prompt"].get("value") == "True":314 prompt = json.dumps(315 {316 "conversation_id": "sambaverse-conversation-id",317 "messages": [318 {319 "message_id": None,320 "role": "user",321 "content": prompt[0],322 }323 ],324 }325 )326 else:327 prompt = prompt[0]328 else:329 prompt = prompt[0]330 if streaming:331 data = {"instance": prompt, "params": params}332 else:333 data = {"instances": [prompt], "params": params}334 headers = {"key": self.sambastudio_api_key.get_secret_value()}335 336 else:337 raise ValueError(338 f"Unsupported URL{self.sambastudio_url}"339 "only openai, generic v1 and generic v2 APIs are supported"340 )341 342 # make the request to SambaStudio API343 http_session = requests.Session()344 if streaming:345 response = http_session.post(346 self.streaming_url, headers=headers, json=data, stream=True347 )348 else:349 response = http_session.post(350 self.base_url, headers=headers, json=data, stream=False351 )352 if response.status_code != 200:353 raise RuntimeError(354 f"Sambanova / complete call failed with status code "355 f"{response.status_code}."356 f"{response.text}."357 )358 return response359 360 def _process_response(self, response: Response) -> str:361 """362 Process a non streaming response from the api363 364 Args:365 response: A request Response object366 367 Returns368 completion: a string with model generation369 """370 371 # Extract json payload form response372 try:373 response_dict = response.json()374 except Exception as e:375 raise RuntimeError(376 f"Sambanova /complete call failed couldn't get JSON response {e}"377 f"response: {response.text}"378 )379 380 # process response payload for openai compatible API381 if "chat/completions" in self.sambastudio_url:382 completion = response_dict["choices"][0]["message"]["content"]383 # process response payload for generic v2 API384 elif "api/v2/predict/generic" in self.sambastudio_url:385 completion = response_dict["items"][0]["value"]["completion"]386 # process response payload for generic v1 API387 elif "api/predict/generic" in self.sambastudio_url:388 completion = response_dict["predictions"][0]["completion"]389 else:390 raise ValueError(391 f"Unsupported URL{self.sambastudio_url}"392 "only openai, generic v1 and generic v2 APIs are supported"393 )394 return completion395 396 def _process_stream_response(self, response: Response) -> Iterator[GenerationChunk]:397 """398 Process a streaming response from the api399 400 Args:401 response: An iterable request Response object402 403 Yields:404 GenerationChunk: a GenerationChunk with model partial generation405 """406 407 try:408 import sseclient409 except ImportError:410 raise ImportError(411 "could not import sseclient library"412 "Please install it with `pip install sseclient-py`."413 )414 415 # process response payload for openai compatible API416 if "chat/completions" in self.sambastudio_url:417 client = sseclient.SSEClient(response)418 for event in client.events():419 if event.event == "error_event":420 raise RuntimeError(421 f"Sambanova /complete call failed with status code "422 f"{response.status_code}."423 f"{event.data}."424 )425 try:426 # check if the response is not a final event ("[DONE]")427 if event.data != "[DONE]":428 if isinstance(event.data, str):429 data = json.loads(event.data)430 else:431 raise RuntimeError(432 f"Sambanova /complete call failed with status code "433 f"{response.status_code}."434 f"{event.data}."435 )436 if data.get("error"):437 raise RuntimeError(438 f"Sambanova /complete call failed with status code "439 f"{response.status_code}."440 f"{event.data}."441 )442 if len(data["choices"]) > 0:443 content = data["choices"][0]["delta"]["content"]444 else:445 content = ""446 generated_chunk = GenerationChunk(text=content)447 yield generated_chunk448 449 except Exception as e:450 raise RuntimeError(451 f"Error getting content chunk raw streamed response: {e}"452 f"data: {event.data}"453 )454 455 # process response payload for generic v2 API456 elif "api/v2/predict/generic" in self.sambastudio_url:457 for line in response.iter_lines():458 try:459 data = json.loads(line)460 content = data["result"]["items"][0]["value"]["stream_token"]461 generated_chunk = GenerationChunk(text=content)462 yield generated_chunk463 464 except Exception as e:465 raise RuntimeError(466 f"Error getting content chunk raw streamed response: {e}"467 f"line: {line}"468 )469 470 # process response payload for generic v1 API471 elif "api/predict/generic" in self.sambastudio_url:472 for line in response.iter_lines():473 try:474 data = json.loads(line)475 content = data["result"]["responses"][0]["stream_token"]476 generated_chunk = GenerationChunk(text=content)477 yield generated_chunk478 479 except Exception as e:480 raise RuntimeError(481 f"Error getting content chunk raw streamed response: {e}"482 f"line: {line}"483 )484 485 else:486 raise ValueError(487 f"Unsupported URL{self.sambastudio_url}"488 "only openai, generic v1 and generic v2 APIs are supported"489 )490 491 def _stream(492 self,493 prompt: Union[List[str], str],494 stop: Optional[List[str]] = None,495 run_manager: Optional[CallbackManagerForLLMRun] = None,496 **kwargs: Any,497 ) -> Iterator[GenerationChunk]:498 """Call out to Sambanova's complete endpoint.499 500 Args:501 prompt: The prompt to pass into the model.502 stop: a list of strings on which the model should stop generating.503 run_manager: A run manager with callbacks for the LLM.504 Yields:505 chunk: GenerationChunk with model partial generation506 """507 response = self._handle_request(prompt, stop, streaming=True)508 for chunk in self._process_stream_response(response):509 if run_manager:510 run_manager.on_llm_new_token(chunk.text)511 yield chunk512 513 def _call(514 self,515 prompt: Union[List[str], str],516 stop: Optional[List[str]] = None,517 run_manager: Optional[CallbackManagerForLLMRun] = None,518 **kwargs: Any,519 ) -> str:520 """Call out to Sambanova's complete endpoint.521 522 Args:523 prompt: The prompt to pass into the model.524 stop: a list of strings on which the model should stop generating.525 526 Returns:527 result: string with model generation528 """529 if self.streaming:530 completion = ""531 for chunk in self._stream(532 prompt=prompt, stop=stop, run_manager=run_manager, **kwargs533 ):534 completion += chunk.text535 536 return completion537 538 response = self._handle_request(prompt, stop, streaming=False)539 completion = self._process_response(response)540 return completion541 542 543class SambaNovaCloud(LLM):544 """545 SambaNova Cloud large language models.546 547 Setup:548 To use, you should have the environment variables:549 ``SAMBANOVA_URL`` set with SambaNova Cloud URL.550 defaults to http://cloud.sambanova.ai/551 ``SAMBANOVA_API_KEY`` set with your SambaNova Cloud API Key.552 Example:553 .. code-block:: python554 from langchain_community.llms.sambanova import SambaNovaCloud555 SambaNovaCloud(556 sambanova_api_key="your-SambaNovaCloud-API-key,557 model = model name,558 max_tokens = max number of tokens to generate,559 temperature = model temperature,560 top_p = model top p,561 top_k = model top k562 )563 Key init args — completion params:564 model: str565 The name of the model to use, e.g., Meta-Llama-3-70B-Instruct-4096566 (set for CoE endpoints).567 streaming: bool568 Whether to use streaming handler when using non streaming methods569 max_tokens: int570 max tokens to generate571 temperature: float572 model temperature573 top_p: float574 model top p575 top_k: int576 model top k577 578 Key init args — client params:579 sambanova_url: str580 SambaNovaCloud Url defaults to http://cloud.sambanova.ai/581 sambanova_api_key: str582 SambaNovaCloud api key583 Instantiate:584 .. code-block:: python585 from langchain_community.llms.sambanova import SambaNovaCloud586 SambaNovaCloud(587 sambanova_api_key="your-SambaNovaCloud-API-key,588 model = model name,589 max_tokens = max number of tokens to generate,590 temperature = model temperature,591 top_p = model top p,592 top_k = model top k593 )594 Invoke:595 .. code-block:: python596 prompt = "tell me a joke"597 response = llm.invoke(prompt)598 Stream:599 .. code-block:: python600 for chunk in llm.stream(prompt):601 print(chunk, end="", flush=True)602 Async:603 .. code-block:: python604 response = llm.ainvoke(prompt)605 await response606 """607 608 sambanova_url: str = Field(default="")609 """SambaNova Cloud Url"""610 611 sambanova_api_key: SecretStr = Field(default=SecretStr(""))612 """SambaNova Cloud api key"""613 614 model: str = Field(default="Meta-Llama-3.1-8B-Instruct")615 """The name of the model"""616 617 streaming: bool = Field(default=False)618 """Whether to use streaming handler when using non streaming methods"""619 620 max_tokens: int = Field(default=1024)621 """max tokens to generate"""622 623 temperature: float = Field(default=0.7)624 """model temperature"""625 626 top_p: Optional[float] = Field(default=None)627 """model top p"""628 629 top_k: Optional[int] = Field(default=None)630 """model top k"""631 632 stream_options: dict = Field(default={"include_usage": True})633 """stream options, include usage to get generation metrics"""634 635 class Config:636 populate_by_name = True637 638 @classmethod639 def is_lc_serializable(cls) -> bool:640 """Return whether this model can be serialized by Langchain."""641 return False642 643 @property644 def lc_secrets(self) -> Dict[str, str]:645 return {"sambanova_api_key": "sambanova_api_key"}646 647 @property648 def _identifying_params(self) -> Dict[str, Any]:649 """Return a dictionary of identifying parameters.650 651 This information is used by the LangChain callback system, which652 is used for tracing purposes make it possible to monitor LLMs.653 """654 return {655 "model": self.model,656 "streaming": self.streaming,657 "max_tokens": self.max_tokens,658 "temperature": self.temperature,659 "top_p": self.top_p,660 "top_k": self.top_k,661 "stream_options": self.stream_options,662 }663 664 @property665 def _llm_type(self) -> str:666 """Get the type of language model used by this chat model."""667 return "sambanovacloud-llm"668 669 def __init__(self, **kwargs: Any) -> None:670 """init and validate environment variables"""671 kwargs["sambanova_url"] = get_from_dict_or_env(672 kwargs,673 "sambanova_url",674 "SAMBANOVA_URL",675 default="https://api.sambanova.ai/v1/chat/completions",676 )677 kwargs["sambanova_api_key"] = convert_to_secret_str(678 get_from_dict_or_env(kwargs, "sambanova_api_key", "SAMBANOVA_API_KEY")679 )680 super().__init__(**kwargs)681 682 def _handle_request(683 self,684 prompt: Union[List[str], str],685 stop: Optional[List[str]] = None,686 streaming: Optional[bool] = False,687 ) -> Response:688 """689 Performs a post request to the LLM API.690 691 Args:692 prompt: The prompt to pass into the model.693 stop: list of stop tokens694 695 Returns:696 A request Response object697 """698 if isinstance(prompt, str):699 prompt = [prompt]700 701 messages_dict = [{"role": "user", "content": prompt[0]}]702 data = {703 "messages": messages_dict,704 "stream": streaming,705 "max_tokens": self.max_tokens,706 "stop": stop,707 "model": self.model,708 "temperature": self.temperature,709 "top_p": self.top_p,710 "top_k": self.top_k,711 }712 data = {key: value for key, value in data.items() if value is not None}713 headers = {714 "Authorization": f"Bearer {self.sambanova_api_key.get_secret_value()}",715 "Content-Type": "application/json",716 }717 718 http_session = requests.Session()719 if streaming:720 response = http_session.post(721 self.sambanova_url, headers=headers, json=data, stream=True722 )723 else:724 response = http_session.post(725 self.sambanova_url, headers=headers, json=data, stream=False726 )727 728 if response.status_code != 200:729 raise RuntimeError(730 f"Sambanova / complete call failed with status code "731 f"{response.status_code}."732 f"{response.text}."733 )734 return response735 736 def _process_response(self, response: Response) -> str:737 """738 Process a non streaming response from the api739 740 Args:741 response: A request Response object742 743 Returns744 completion: a string with model generation745 """746 747 # Extract json payload form response748 try:749 response_dict = response.json()750 except Exception as e:751 raise RuntimeError(752 f"Sambanova /complete call failed couldn't get JSON response {e}"753 f"response: {response.text}"754 )755 756 completion = response_dict["choices"][0]["message"]["content"]757 758 return completion759 760 def _process_stream_response(self, response: Response) -> Iterator[GenerationChunk]:761 """762 Process a streaming response from the api763 764 Args:765 response: An iterable request Response object766 767 Yields:768 GenerationChunk: a GenerationChunk with model partial generation769 """770 771 try:772 import sseclient773 except ImportError:774 raise ImportError(775 "could not import sseclient library"776 "Please install it with `pip install sseclient-py`."777 )778 779 client = sseclient.SSEClient(response)780 for event in client.events():781 if event.event == "error_event":782 raise RuntimeError(783 f"Sambanova /complete call failed with status code "784 f"{response.status_code}."785 f"{event.data}."786 )787 try:788 # check if the response is not a final event ("[DONE]")789 if event.data != "[DONE]":790 if isinstance(event.data, str):791 data = json.loads(event.data)792 else:793 raise RuntimeError(794 f"Sambanova /complete call failed with status code "795 f"{response.status_code}."796 f"{event.data}."797 )798 if data.get("error"):799 raise RuntimeError(800 f"Sambanova /complete call failed with status code "801 f"{response.status_code}."802 f"{event.data}."803 )804 if len(data["choices"]) > 0:805 content = data["choices"][0]["delta"]["content"]806 else:807 content = ""808 generated_chunk = GenerationChunk(text=content)809 yield generated_chunk810 811 except Exception as e:812 raise RuntimeError(813 f"Error getting content chunk raw streamed response: {e}"814 f"data: {event.data}"815 )816 817 def _call(818 self,819 prompt: Union[List[str], str],820 stop: Optional[List[str]] = None,821 run_manager: Optional[CallbackManagerForLLMRun] = None,822 **kwargs: Any,823 ) -> str:824 """Call out to SambaNovaCloud complete endpoint.825 826 Args:827 prompt: The prompt to pass into the model.828 stop: Optional list of stop words to use when generating.829 830 Returns:831 The string generated by the model.832 """833 if self.streaming:834 completion = ""835 for chunk in self._stream(836 prompt=prompt, stop=stop, run_manager=run_manager, **kwargs837 ):838 completion += chunk.text839 840 return completion841 842 response = self._handle_request(prompt, stop, streaming=False)843 completion = self._process_response(response)844 return completion845 846 def _stream(847 self,848 prompt: Union[List[str], str],849 stop: Optional[List[str]] = None,850 run_manager: Optional[CallbackManagerForLLMRun] = None,851 **kwargs: Any,852 ) -> Iterator[GenerationChunk]:853 """Call out to SambaNovaCloud complete endpoint.854 855 Args:856 prompt: The prompt to pass into the model.857 stop: Optional list of stop words to use when generating.858 859 Returns:860 The string generated by the model.861 """862 response = self._handle_request(prompt, stop, streaming=True)863 for chunk in self._process_stream_response(response):864 if run_manager:865 run_manager.on_llm_new_token(chunk.text)866 yield chunk867 