codekingpro/portable-devtools
114k
1"""Wrapper around Together AI's Completion API."""2 3import logging4from typing import Any, Dict, List, Optional5 6from aiohttp import ClientSession7from langchain_core._api.deprecation import deprecated8from langchain_core.callbacks import (9 AsyncCallbackManagerForLLMRun,10 CallbackManagerForLLMRun,11)12from langchain_core.language_models.llms import LLM13from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env14from pydantic import ConfigDict, SecretStr, model_validator15 16from langchain_community.utilities.requests import Requests17 18logger = logging.getLogger(__name__)19 20 21@deprecated(22 since="0.0.12", removal="1.0", alternative_import="langchain_together.Together"23)24class Together(LLM):25 """LLM models from `Together`.26 27 To use, you'll need an API key which you can find here:28 https://api.together.xyz/settings/api-keys. This can be passed in as init param29 ``together_api_key`` or set as environment variable ``TOGETHER_API_KEY``.30 31 Together AI API reference: https://docs.together.ai/reference/inference32 """33 34 base_url: str = "https://api.together.xyz/inference"35 """Base inference API URL."""36 together_api_key: SecretStr37 """Together AI API key. Get it here: https://api.together.xyz/settings/api-keys"""38 model: str39 """Model name. Available models listed here: 40 https://docs.together.ai/docs/inference-models41 """42 temperature: Optional[float] = None43 """Model temperature."""44 top_p: Optional[float] = None45 """Used to dynamically adjust the number of choices for each predicted token based 46 on the cumulative probabilities. A value of 1 will always yield the same 47 output. A temperature less than 1 favors more correctness and is appropriate 48 for question answering or summarization. A value greater than 1 introduces more 49 randomness in the output.50 """51 top_k: Optional[int] = None52 """Used to limit the number of choices for the next predicted word or token. It 53 specifies the maximum number of tokens to consider at each step, based on their 54 probability of occurrence. This technique helps to speed up the generation 55 process and can improve the quality of the generated text by focusing on the 56 most likely options.57 """58 max_tokens: Optional[int] = None59 """The maximum number of tokens to generate."""60 repetition_penalty: Optional[float] = None61 """A number that controls the diversity of generated text by reducing the 62 likelihood of repeated sequences. Higher values decrease repetition.63 """64 logprobs: Optional[int] = None65 """An integer that specifies how many top token log probabilities are included in 66 the response for each token generation step.67 """68 69 model_config = ConfigDict(70 extra="forbid",71 )72 73 @model_validator(mode="before")74 @classmethod75 def validate_environment(cls, values: Dict) -> Any:76 """Validate that api key exists in environment."""77 values["together_api_key"] = convert_to_secret_str(78 get_from_dict_or_env(values, "together_api_key", "TOGETHER_API_KEY")79 )80 return values81 82 @property83 def _llm_type(self) -> str:84 """Return type of model."""85 return "together"86 87 def _format_output(self, output: dict) -> str:88 return output["output"]["choices"][0]["text"]89 90 @staticmethod91 def get_user_agent() -> str:92 from langchain_community import __version__93 94 return f"langchain/{__version__}"95 96 @property97 def default_params(self) -> Dict[str, Any]:98 return {99 "model": self.model,100 "temperature": self.temperature,101 "top_p": self.top_p,102 "top_k": self.top_k,103 "max_tokens": self.max_tokens,104 "repetition_penalty": self.repetition_penalty,105 }106 107 def _call(108 self,109 prompt: str,110 stop: Optional[List[str]] = None,111 run_manager: Optional[CallbackManagerForLLMRun] = None,112 **kwargs: Any,113 ) -> str:114 """Call out to Together's text generation endpoint.115 116 Args:117 prompt: The prompt to pass into the model.118 119 Returns:120 The string generated by the model..121 """122 123 headers = {124 "Authorization": f"Bearer {self.together_api_key.get_secret_value()}",125 "Content-Type": "application/json",126 }127 stop_to_use = stop[0] if stop and len(stop) == 1 else stop128 payload: Dict[str, Any] = {129 **self.default_params,130 "prompt": prompt,131 "stop": stop_to_use,132 **kwargs,133 }134 135 # filter None values to not pass them to the http payload136 payload = {k: v for k, v in payload.items() if v is not None}137 request = Requests(headers=headers)138 response = request.post(url=self.base_url, data=payload)139 140 if response.status_code >= 500:141 raise Exception(f"Together Server: Error {response.status_code}")142 elif response.status_code >= 400:143 raise ValueError(f"Together received an invalid payload: {response.text}")144 elif response.status_code != 200:145 raise Exception(146 f"Together returned an unexpected response with status "147 f"{response.status_code}: {response.text}"148 )149 150 data = response.json()151 if data.get("status") != "finished":152 err_msg = data.get("error", "Undefined Error")153 raise Exception(err_msg)154 155 output = self._format_output(data)156 157 return output158 159 async def _acall(160 self,161 prompt: str,162 stop: Optional[List[str]] = None,163 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,164 **kwargs: Any,165 ) -> str:166 """Call Together model to get predictions based on the prompt.167 168 Args:169 prompt: The prompt to pass into the model.170 171 Returns:172 The string generated by the model.173 """174 headers = {175 "Authorization": f"Bearer {self.together_api_key.get_secret_value()}",176 "Content-Type": "application/json",177 }178 stop_to_use = stop[0] if stop and len(stop) == 1 else stop179 payload: Dict[str, Any] = {180 **self.default_params,181 "prompt": prompt,182 "stop": stop_to_use,183 **kwargs,184 }185 186 # filter None values to not pass them to the http payload187 payload = {k: v for k, v in payload.items() if v is not None}188 async with ClientSession() as session:189 async with session.post(190 self.base_url, json=payload, headers=headers191 ) as response:192 if response.status >= 500:193 raise Exception(f"Together Server: Error {response.status}")194 elif response.status >= 400:195 raise ValueError(196 f"Together received an invalid payload: {response.text}"197 )198 elif response.status != 200:199 raise Exception(200 f"Together returned an unexpected response with status "201 f"{response.status}: {response.text}"202 )203 204 response_json = await response.json()205 206 if response_json.get("status") != "finished":207 err_msg = response_json.get("error", "Undefined Error")208 raise Exception(err_msg)209 210 output = self._format_output(response_json)211 return output212 