diff --git a/libs/langchain/langchain/chat_models/fireworks.py b/libs/langchain/langchain/chat_models/fireworks.py index fd851c0eedc70..e63202183b912 100644 --- a/libs/langchain/langchain/chat_models/fireworks.py +++ b/libs/langchain/langchain/chat_models/fireworks.py @@ -1,3 +1,5 @@ +import asyncio +from concurrent.futures import ThreadPoolExecutor from typing import ( Any, AsyncIterator, @@ -11,12 +13,16 @@ ) from langchain.adapters.openai import convert_message_to_dict +from langchain.callbacks.base import Callbacks from langchain.callbacks.manager import ( + AsyncCallbackManager, AsyncCallbackManagerForLLMRun, + CallbackManager, CallbackManagerForLLMRun, ) from langchain.chat_models.base import BaseChatModel from langchain.llms.base import create_base_retry_decorator +from langchain.load.dump import dumpd from langchain.pydantic_v1 import Field, root_validator from langchain.schema.messages import ( AIMessage, @@ -32,7 +38,13 @@ SystemMessage, SystemMessageChunk, ) -from langchain.schema.output import ChatGeneration, ChatGenerationChunk, ChatResult +from langchain.schema.output import ( + ChatGeneration, + ChatGenerationChunk, + ChatResult, + LLMResult, + RunInfo, +) from langchain.utils.env import get_from_dict_or_env @@ -89,6 +101,7 @@ class ChatFireworks(BaseChatModel): ) fireworks_api_key: Optional[str] = None max_retries: int = 20 + batch_size: int = 20 @root_validator() def validate_environment(cls, values: Dict) -> Dict: @@ -108,6 +121,194 @@ def _llm_type(self) -> str: """Return type of llm.""" return "fireworks-chat" + def generate( + self, + messages: List[List[BaseMessage]], + stop: Optional[List[str]] = None, + callbacks: Callbacks = None, + *, + tags: Optional[List[str]] = None, + metadata: Optional[Dict[str, Any]] = None, + run_name: Optional[str] = None, + **kwargs: Any, + ) -> LLMResult: + """Top Level call""" + params = self._get_invocation_params(stop=stop, **kwargs) + options = {"stop": stop} + + callback_manager = CallbackManager.configure( + callbacks, + self.callbacks, + self.verbose, + tags, + self.tags, + metadata, + self.metadata, + ) + run_managers = callback_manager.on_chat_model_start( + dumpd(self), + messages, + invocation_params=params, + options=options, + name=run_name, + ) + + def _completion_with_retry_batching(message): + args_list = [ + (m, stop, run_managers[i] if run_managers else None, kwargs) + for i, m in enumerate(message) + ] + with ThreadPoolExecutor() as executor: + results = list(executor.map(self._process_message, args_list)) + + return results + + sub_messages = self.get_batch_messages(params, messages, stop) + + results = [] + for message in sub_messages: + results.extend(_completion_with_retry_batching(message)) + + flattened_outputs = [ + LLMResult(generations=[res.generations], llm_output=res.llm_output) + for res in results + ] + llm_output = self._combine_llm_outputs([res.llm_output for res in results]) + generations = [res.generations for res in results] + output = LLMResult(generations=generations, llm_output=llm_output) + if run_managers: + run_infos = [] + for manager, flattened_output in zip(run_managers, flattened_outputs): + manager.on_llm_end(flattened_output) + run_infos.append(RunInfo(run_id=manager.run_id)) + output.run = run_infos + return output + + async def agenerate( + self, + messages: List[List[BaseMessage]], + stop: Optional[List[str]] = None, + callbacks: Callbacks = None, + *, + tags: Optional[List[str]] = None, + metadata: Optional[Dict[str, Any]] = None, + run_name: Optional[str] = None, + **kwargs: Any, + ) -> LLMResult: + """Top Level call""" + params = self._get_invocation_params(stop=stop, **kwargs) + options = {"stop": stop} + + callback_manager = AsyncCallbackManager.configure( + callbacks, + self.callbacks, + self.verbose, + tags, + self.tags, + metadata, + self.metadata, + ) + + run_managers = await callback_manager.on_chat_model_start( + dumpd(self), + messages, + invocation_params=params, + options=options, + name=run_name, + ) + + async def _acompletion_with_retry_batching(message): + args_list = [ + (m, stop, run_managers[i] if run_managers else None, kwargs) + for i, m in enumerate(messages) + ] + loop = asyncio.get_event_loop() + with ThreadPoolExecutor() as executor: + results = await asyncio.gather( + *[ + loop.run_in_executor(executor, self._process_message, args) + for args in args_list + ], + return_exceptions=True, + ) + + return results + + sub_messages = self.get_batch_messages(params, messages, stop) + + results = [] + for message in sub_messages: + results.extend(await _acompletion_with_retry_batching(message)) + + exceptions = [] + for i, res in enumerate(results): + if isinstance(res, BaseException): + if run_managers: + await run_managers[i].on_llm_error(res) + exceptions.append(res) + if exceptions: + if run_managers: + await asyncio.gather( + *[ + run_manager.on_llm_end( + LLMResult( + generations=[res.generations], llm_output=res.llm_output + ) + ) + for run_manager, res in zip(run_managers, results) + if not isinstance(res, Exception) + ] + ) + raise exceptions[0] + flattened_outputs = [ + LLMResult(generations=[res.generations], llm_output=res.llm_output) + for res in results + ] + llm_output = self._combine_llm_outputs([res.llm_output for res in results]) + generations = [res.generations for res in results] + output = LLMResult(generations=generations, llm_output=llm_output) + await asyncio.gather( + *[ + run_manager.on_llm_end(flattened_output) + for run_manager, flattened_output in zip( + run_managers, flattened_outputs + ) + ] + ) + if run_managers: + output.run = [ + RunInfo(run_id=run_manager.run_id) for run_manager in run_managers + ] + return output + + def _process_message(self, args): + m, stop, run_manager, kwargs = args + try: + return self._generate_with_cache( + m, stop=stop, run_manager=run_manager, **kwargs + ) + except BaseException as e: + if run_manager: + run_manager.on_llm_error(e) + raise e + + def get_batch_messages( + self, + params: Dict[str, Any], + messages: List[List[BaseMessage]], + stop: Optional[List[str]] = None, + ) -> List[List[str]]: + """Get the sub messages for llm call.""" + if stop is not None: + if "stop" in params: + raise ValueError("`stop` found in both the input and default params.") + + sub_messages = [ + messages[i : i + self.batch_size] + for i in range(0, len(messages), self.batch_size) + ] + return sub_messages + def _generate( self, messages: List[BaseMessage], diff --git a/libs/langchain/langchain/llms/fireworks.py b/libs/langchain/langchain/llms/fireworks.py index 6922b2a6e92b7..406a8d28e5a4c 100644 --- a/libs/langchain/langchain/llms/fireworks.py +++ b/libs/langchain/langchain/llms/fireworks.py @@ -1,4 +1,15 @@ -from typing import Any, AsyncIterator, Callable, Dict, Iterator, List, Optional, Union +import asyncio +from concurrent.futures import ThreadPoolExecutor +from typing import ( + Any, + AsyncIterator, + Callable, + Dict, + Iterator, + List, + Optional, + Union, +) from langchain.callbacks.manager import ( AsyncCallbackManagerForLLMRun, @@ -6,9 +17,7 @@ ) from langchain.llms.base import LLM, create_base_retry_decorator from langchain.pydantic_v1 import Field, root_validator -from langchain.schema.language_model import LanguageModelInput -from langchain.schema.output import GenerationChunk -from langchain.schema.runnable.config import RunnableConfig +from langchain.schema.output import Generation, GenerationChunk, LLMResult from langchain.utils.env import get_from_dict_or_env @@ -38,6 +47,7 @@ class Fireworks(LLM): ) fireworks_api_key: Optional[str] = None max_retries: int = 20 + batch_size: int = 20 @root_validator() def validate_environment(cls, values: Dict) -> Dict: @@ -95,6 +105,87 @@ async def _acall( return response.choices[0].text + def _generate( + self, + prompts: List[str], + stop: Optional[List[str]] = None, + run_manager: Optional[CallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> LLMResult: + """Call out to Fireworks endpoint with k unique prompts. + Args: + prompts: The prompts to pass into the model. + stop: Optional list of stop words to use when generating. + Returns: + The full LLM output. + """ + params = { + "model": self.model, + **self.model_kwargs, + } + sub_prompts = self.get_batch_prompts(params, prompts, stop) + choices = [] + for _prompts in sub_prompts: + response = completion_with_retry_batching(self, prompt=_prompts, **params) + choices.extend(response) + + return self.create_llm_result(choices, prompts) + + async def _agenerate( + self, + prompts: List[str], + stop: Optional[List[str]] = None, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> LLMResult: + """Call out to Fireworks endpoint async with k unique prompts.""" + params = { + "model": self.model, + **self.model_kwargs, + } + sub_prompts = self.get_batch_prompts(params, prompts, stop) + choices = [] + for _prompts in sub_prompts: + response = await acompletion_with_retry_batching( + self, prompt=_prompts, **params + ) + choices.extend(response) + + return self.create_llm_result(choices, prompts) + + def get_batch_prompts( + self, + params: Dict[str, Any], + prompts: List[str], + stop: Optional[List[str]] = None, + ) -> List[List[str]]: + """Get the sub prompts for llm call.""" + if stop is not None: + if "stop" in params: + raise ValueError("`stop` found in both the input and default params.") + + sub_prompts = [ + prompts[i : i + self.batch_size] + for i in range(0, len(prompts), self.batch_size) + ] + return sub_prompts + + def create_llm_result(self, choices: Any, prompts: List[str]) -> LLMResult: + """Create the LLMResult from the choices and prompts.""" + generations = [] + for i, _ in enumerate(prompts): + sub_choices = choices[i : (i + 1)] + generations.append( + [ + Generation( + text=choice.__dict__["choices"][0].text, + ) + for choice in sub_choices + ] + ) + llm_output = {"model": self.model} + return LLMResult(generations=generations, llm_output=llm_output) + def _stream( self, prompt: str, @@ -108,7 +199,7 @@ def _stream( "stream": True, **self.model_kwargs, } - for stream_resp in completion_with_retry( + for stream_resp in completion_with_retry_streaming( self, run_manager=run_manager, stop=stop, **params ): chunk = _stream_response_to_generation_chunk(stream_resp) @@ -133,42 +224,6 @@ async def _astream( chunk = _stream_response_to_generation_chunk(stream_resp) yield chunk - def stream( - self, - input: LanguageModelInput, - config: Optional[RunnableConfig] = None, - *, - stop: Optional[List[str]] = None, - **kwargs: Any, - ) -> Iterator[str]: - prompt = self._convert_input(input).to_string() - generation: Optional[GenerationChunk] = None - for chunk in self._stream(prompt): - yield chunk.text - if generation is None: - generation = chunk - else: - generation += chunk - assert generation is not None - - async def astream( - self, - input: LanguageModelInput, - config: Optional[RunnableConfig] = None, - *, - stop: Optional[List[str]] = None, - **kwargs: Any, - ) -> AsyncIterator[str]: - prompt = self._convert_input(input).to_string() - generation: Optional[GenerationChunk] = None - async for chunk in self._astream(prompt): - yield chunk.text - if generation is None: - generation = chunk - else: - generation += chunk - assert generation is not None - def completion_with_retry( llm: Fireworks, @@ -210,6 +265,91 @@ async def _completion_with_retry(**kwargs: Any) -> Any: return await _completion_with_retry(**kwargs) +def completion_with_retry_batching( + llm: Fireworks, + *, + run_manager: Optional[CallbackManagerForLLMRun] = None, + **kwargs: Any, +) -> Any: + """Use tenacity to retry the completion call.""" + import fireworks.client + + prompt = kwargs["prompt"] + del kwargs["prompt"] + + retry_decorator = _create_retry_decorator(llm, run_manager=run_manager) + + @retry_decorator + def _completion_with_retry(prompt) -> Any: + return fireworks.client.Completion.create(**kwargs, prompt=prompt) + + def batch_sync_run(): + with ThreadPoolExecutor() as executor: + results = list(executor.map(_completion_with_retry, prompt)) + return results + + return batch_sync_run() + + +async def acompletion_with_retry_batching( + llm: Fireworks, + *, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, +) -> Any: + """Use tenacity to retry the completion call.""" + import fireworks.client + + prompt = kwargs["prompt"] + del kwargs["prompt"] + + retry_decorator = _create_retry_decorator(llm, run_manager=run_manager) + + @retry_decorator + async def _completion_with_retry(prompt) -> Any: + return await fireworks.client.Completion.acreate(**kwargs, prompt=prompt) + + def run_coroutine_in_new_loop(coroutine_func, *args, **kwargs): + new_loop = asyncio.new_event_loop() + try: + asyncio.set_event_loop(new_loop) + return new_loop.run_until_complete(coroutine_func(*args, **kwargs)) + finally: + new_loop.close() + + async def batch_sync_run(): + with ThreadPoolExecutor() as executor: + results = list( + executor.map( + run_coroutine_in_new_loop, + [_completion_with_retry] * len(prompt), + prompt, + ) + ) + return results + + return await batch_sync_run() + + +def completion_with_retry_streaming( + llm: Fireworks, + *, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, +) -> Any: + import fireworks.client + + retry_decorator = _create_retry_decorator(llm, run_manager=run_manager) + + @retry_decorator + def _completion_with_retry(**kwargs: Any) -> Any: + return fireworks.client.Completion.create( + **kwargs, + ) + + return _completion_with_retry(**kwargs) + + async def acompletion_with_retry_streaming( llm: Fireworks, *, diff --git a/libs/langchain/tests/integration_tests/llms/test_fireworks.py b/libs/langchain/tests/integration_tests/llms/test_fireworks.py index cbc7473665a4b..2c1585769b4ad 100644 --- a/libs/langchain/tests/integration_tests/llms/test_fireworks.py +++ b/libs/langchain/tests/integration_tests/llms/test_fireworks.py @@ -86,3 +86,24 @@ async def test_fireworks_multiple_prompts_async_agenerate() -> None: assert isinstance(output, LLMResult) assert isinstance(output.generations, list) assert len(output.generations) == 2 + + +def test_fireworks_batch() -> None: + """Test streaming tokens from Fireworks.""" + llm = Fireworks() + + result = llm.batch(["How is the weather in New York today?", "I'm pickle rick"]) + for token in result: + assert isinstance(token, str) + + +@pytest.mark.asyncio +async def test_fireworks_abatch() -> None: + """Test streaming tokens from Fireworks.""" + llm = Fireworks() + + result = await llm.abatch( + ["How is the weather in New York today?", "I'm not Pickle Rick"] + ) + for token in result: + assert isinstance(token, str)