Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fireworks.py388 linesDownload Raw Back to llms
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 
codekingpro/portable-devtools · Team Ai