codekingpro/portable-devtools
114k
1import asyncio2from concurrent.futures import ThreadPoolExecutor3from typing import Any, AsyncIterator, Callable, Dict, Iterator, List, Optional, Union4 5from langchain_core._api.deprecation import deprecated6from langchain_core.callbacks import (7 AsyncCallbackManagerForLLMRun,8 CallbackManagerForLLMRun,9)10from langchain_core.language_models.llms import BaseLLM, create_base_retry_decorator11from langchain_core.outputs import Generation, GenerationChunk, LLMResult12from langchain_core.utils import convert_to_secret_str, pre_init13from langchain_core.utils.env import get_from_dict_or_env14from pydantic import Field, SecretStr15 16 17def _stream_response_to_generation_chunk(18 stream_response: Any,19) -> GenerationChunk:20 """Convert a stream response to a generation chunk."""21 return GenerationChunk(22 text=stream_response.choices[0].text,23 generation_info=dict(24 finish_reason=stream_response.choices[0].finish_reason,25 logprobs=stream_response.choices[0].logprobs,26 ),27 )28 29 30@deprecated(31 since="0.0.26",32 removal="1.0",33 alternative_import="langchain_fireworks.Fireworks",34)35class Fireworks(BaseLLM):36 """Fireworks models."""37 38 model: str = "accounts/fireworks/models/llama-v2-7b-chat"39 model_kwargs: dict = Field(40 default_factory=lambda: {41 "temperature": 0.7,42 "max_tokens": 512,43 "top_p": 1,44 }.copy()45 )46 fireworks_api_key: Optional[SecretStr] = None47 max_retries: int = 2048 batch_size: int = 2049 use_retry: bool = True50 51 @property52 def lc_secrets(self) -> Dict[str, str]:53 return {"fireworks_api_key": "FIREWORKS_API_KEY"}54 55 @classmethod56 def is_lc_serializable(cls) -> bool:57 return True58 59 @classmethod60 def get_lc_namespace(cls) -> List[str]:61 """Get the namespace of the langchain object."""62 return ["langchain", "llms", "fireworks"]63 64 @pre_init65 def validate_environment(cls, values: Dict) -> Dict:66 """Validate that api key in environment."""67 try:68 import fireworks.client69 except ImportError as e:70 raise ImportError(71 "Could not import fireworks-ai python package. "72 "Please install it with `pip install fireworks-ai`."73 ) from e74 fireworks_api_key = convert_to_secret_str(75 get_from_dict_or_env(values, "fireworks_api_key", "FIREWORKS_API_KEY")76 )77 fireworks.client.api_key = fireworks_api_key.get_secret_value()78 return values79 80 @property81 def _llm_type(self) -> str:82 """Return type of llm."""83 return "fireworks"84 85 def _generate(86 self,87 prompts: List[str],88 stop: Optional[List[str]] = None,89 run_manager: Optional[CallbackManagerForLLMRun] = None,90 **kwargs: Any,91 ) -> LLMResult:92 """Call out to Fireworks endpoint with k unique prompts.93 Args:94 prompts: The prompts to pass into the model.95 stop: Optional list of stop words to use when generating.96 Returns:97 The full LLM output.98 """99 params = {100 "model": self.model,101 **self.model_kwargs,102 }103 sub_prompts = self.get_batch_prompts(prompts)104 choices = []105 for _prompts in sub_prompts:106 response = completion_with_retry_batching(107 self,108 self.use_retry,109 prompt=_prompts,110 run_manager=run_manager,111 stop=stop,112 **params,113 )114 choices.extend(response)115 116 return self.create_llm_result(choices, prompts)117 118 async def _agenerate(119 self,120 prompts: List[str],121 stop: Optional[List[str]] = None,122 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,123 **kwargs: Any,124 ) -> LLMResult:125 """Call out to Fireworks endpoint async with k unique prompts."""126 params = {127 "model": self.model,128 **self.model_kwargs,129 }130 sub_prompts = self.get_batch_prompts(prompts)131 choices = []132 for _prompts in sub_prompts:133 response = await acompletion_with_retry_batching(134 self,135 self.use_retry,136 prompt=_prompts,137 run_manager=run_manager,138 stop=stop,139 **params,140 )141 choices.extend(response)142 143 return self.create_llm_result(choices, prompts)144 145 def get_batch_prompts(146 self,147 prompts: List[str],148 ) -> List[List[str]]:149 """Get the sub prompts for llm call."""150 sub_prompts = [151 prompts[i : i + self.batch_size]152 for i in range(0, len(prompts), self.batch_size)153 ]154 return sub_prompts155 156 def create_llm_result(self, choices: Any, prompts: List[str]) -> LLMResult:157 """Create the LLMResult from the choices and prompts."""158 generations = []159 for i, _ in enumerate(prompts):160 sub_choices = choices[i : (i + 1)]161 generations.append(162 [163 Generation(164 text=choice.__dict__["choices"][0].text,165 )166 for choice in sub_choices167 ]168 )169 llm_output = {"model": self.model}170 return LLMResult(generations=generations, llm_output=llm_output)171 172 def _stream(173 self,174 prompt: str,175 stop: Optional[List[str]] = None,176 run_manager: Optional[CallbackManagerForLLMRun] = None,177 **kwargs: Any,178 ) -> Iterator[GenerationChunk]:179 params = {180 "model": self.model,181 "prompt": prompt,182 "stream": True,183 **self.model_kwargs,184 }185 for stream_resp in completion_with_retry(186 self, self.use_retry, run_manager=run_manager, stop=stop, **params187 ):188 chunk = _stream_response_to_generation_chunk(stream_resp)189 if run_manager:190 run_manager.on_llm_new_token(chunk.text, chunk=chunk)191 yield chunk192 193 async def _astream(194 self,195 prompt: str,196 stop: Optional[List[str]] = None,197 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,198 **kwargs: Any,199 ) -> AsyncIterator[GenerationChunk]:200 params = {201 "model": self.model,202 "prompt": prompt,203 "stream": True,204 **self.model_kwargs,205 }206 async for stream_resp in await acompletion_with_retry_streaming(207 self, self.use_retry, run_manager=run_manager, stop=stop, **params208 ):209 chunk = _stream_response_to_generation_chunk(stream_resp)210 if run_manager:211 await run_manager.on_llm_new_token(chunk.text, chunk=chunk)212 yield chunk213 214 215def conditional_decorator(216 condition: bool, decorator: Callable[[Any], Any]217) -> Callable[[Any], Any]:218 """Conditionally apply a decorator.219 220 Args:221 condition: A boolean indicating whether to apply the decorator.222 decorator: A decorator function.223 224 Returns:225 A decorator function.226 """227 228 def actual_decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]:229 if condition:230 return decorator(func)231 return func232 233 return actual_decorator234 235 236def completion_with_retry(237 llm: Fireworks,238 use_retry: bool,239 *,240 run_manager: Optional[CallbackManagerForLLMRun] = None,241 **kwargs: Any,242) -> Any:243 """Use tenacity to retry the completion call."""244 import fireworks.client245 246 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)247 248 @conditional_decorator(use_retry, retry_decorator)249 def _completion_with_retry(**kwargs: Any) -> Any:250 return fireworks.client.Completion.create(251 **kwargs,252 )253 254 return _completion_with_retry(**kwargs)255 256 257async def acompletion_with_retry(258 llm: Fireworks,259 use_retry: bool,260 *,261 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,262 **kwargs: Any,263) -> Any:264 """Use tenacity to retry the completion call."""265 import fireworks.client266 267 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)268 269 @conditional_decorator(use_retry, retry_decorator)270 async def _completion_with_retry(**kwargs: Any) -> Any:271 return await fireworks.client.Completion.acreate(272 **kwargs,273 )274 275 return await _completion_with_retry(**kwargs)276 277 278def completion_with_retry_batching(279 llm: Fireworks,280 use_retry: bool,281 *,282 run_manager: Optional[CallbackManagerForLLMRun] = None,283 **kwargs: Any,284) -> Any:285 """Use tenacity to retry the completion call."""286 import fireworks.client287 288 prompt = kwargs["prompt"]289 del kwargs["prompt"]290 291 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)292 293 @conditional_decorator(use_retry, retry_decorator)294 def _completion_with_retry(prompt: str) -> Any:295 return fireworks.client.Completion.create(**kwargs, prompt=prompt)296 297 def batch_sync_run() -> List:298 with ThreadPoolExecutor() as executor:299 results = list(executor.map(_completion_with_retry, prompt))300 return results301 302 return batch_sync_run()303 304 305async def acompletion_with_retry_batching(306 llm: Fireworks,307 use_retry: bool,308 *,309 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,310 **kwargs: Any,311) -> Any:312 """Use tenacity to retry the completion call."""313 import fireworks.client314 315 prompt = kwargs["prompt"]316 del kwargs["prompt"]317 318 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)319 320 @conditional_decorator(use_retry, retry_decorator)321 async def _completion_with_retry(prompt: str) -> Any:322 return await fireworks.client.Completion.acreate(**kwargs, prompt=prompt)323 324 def run_coroutine_in_new_loop(325 coroutine_func: Any, *args: Dict, **kwargs: Dict326 ) -> Any:327 new_loop = asyncio.new_event_loop()328 try:329 asyncio.set_event_loop(new_loop)330 return new_loop.run_until_complete(coroutine_func(*args, **kwargs))331 finally:332 new_loop.close()333 334 async def batch_sync_run() -> List:335 with ThreadPoolExecutor() as executor:336 results = list(337 executor.map(338 run_coroutine_in_new_loop,339 [_completion_with_retry] * len(prompt),340 prompt,341 )342 )343 return results344 345 return await batch_sync_run()346 347 348async def acompletion_with_retry_streaming(349 llm: Fireworks,350 use_retry: bool,351 *,352 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,353 **kwargs: Any,354) -> Any:355 """Use tenacity to retry the completion call for streaming."""356 import fireworks.client357 358 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)359 360 @conditional_decorator(use_retry, retry_decorator)361 async def _completion_with_retry(**kwargs: Any) -> Any:362 return fireworks.client.Completion.acreate(363 **kwargs,364 )365 366 return await _completion_with_retry(**kwargs)367 368 369def _create_retry_decorator(370 llm: Fireworks,371 *,372 run_manager: Optional[373 Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]374 ] = None,375) -> Callable[[Any], Any]:376 """Define retry mechanism."""377 import fireworks.client378 379 errors = [380 fireworks.client.error.RateLimitError,381 fireworks.client.error.InternalServerError,382 fireworks.client.error.BadGatewayError,383 fireworks.client.error.ServiceUnavailableError,384 ]385 return create_base_retry_decorator(386 error_types=errors, max_retries=llm.max_retries, run_manager=run_manager387 )388 