codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from typing import Any, AsyncIterator, Dict, Iterator, List, Optional4 5from langchain_core.callbacks import (6 AsyncCallbackManagerForLLMRun,7 CallbackManagerForLLMRun,8)9from langchain_core.language_models.chat_models import (10 BaseChatModel,11 agenerate_from_stream,12 generate_from_stream,13)14from langchain_core.messages import (15 AIMessage,16 AIMessageChunk,17 BaseMessage,18 ChatMessage,19 HumanMessage,20 SystemMessage,21)22from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult23 24from langchain_community.llms.friendli import BaseFriendli25 26 27def get_role(message: BaseMessage) -> str:28 """Get role of the message.29 30 Args:31 message (BaseMessage): The message object.32 33 Raises:34 ValueError: Raised when the message is of an unknown type.35 36 Returns:37 str: The role of the message.38 """39 if isinstance(message, ChatMessage) or isinstance(message, HumanMessage):40 return "user"41 if isinstance(message, AIMessage):42 return "assistant"43 if isinstance(message, SystemMessage):44 return "system"45 raise ValueError(f"Got unknown type {message}")46 47 48def get_chat_request(messages: List[BaseMessage]) -> Dict[str, Any]:49 """Get a request of the Friendli chat API.50 51 Args:52 messages (List[BaseMessage]): Messages comprising the conversation so far.53 54 Returns:55 Dict[str, Any]: The request for the Friendli chat API.56 """57 return {58 "messages": [59 {"role": get_role(message), "content": message.content}60 for message in messages61 ]62 }63 64 65class ChatFriendli(BaseChatModel, BaseFriendli):66 """Friendli LLM for chat.67 68 ``friendli-client`` package should be installed with `pip install friendli-client`.69 You must set ``FRIENDLI_TOKEN`` environment variable or provide the value of your70 personal access token for the ``friendli_token`` argument.71 72 Example:73 .. code-block:: python74 75 from langchain_community.chat_models import FriendliChat76 77 chat = Friendli(78 model="meta-llama-3.1-8b-instruct", friendli_token="YOUR FRIENDLI TOKEN"79 )80 chat.invoke("What is generative AI?")81 """82 83 model: str = "meta-llama-3.1-8b-instruct"84 85 @property86 def lc_secrets(self) -> Dict[str, str]:87 return {"friendli_token": "FRIENDLI_TOKEN"}88 89 @property90 def _default_params(self) -> Dict[str, Any]:91 """Get the default parameters for calling Friendli completions API."""92 return {93 "frequency_penalty": self.frequency_penalty,94 "presence_penalty": self.presence_penalty,95 "max_tokens": self.max_tokens,96 "stop": self.stop,97 "temperature": self.temperature,98 "top_p": self.top_p,99 }100 101 @property102 def _identifying_params(self) -> Dict[str, Any]:103 """Get the identifying parameters."""104 return {"model": self.model, **self._default_params}105 106 @property107 def _llm_type(self) -> str:108 return "friendli-chat"109 110 def _get_invocation_params(111 self, stop: Optional[List[str]] = None, **kwargs: Any112 ) -> Dict[str, Any]:113 """Get the parameters used to invoke the model."""114 params = self._default_params115 if self.stop is not None and stop is not None:116 raise ValueError("`stop` found in both the input and default params.")117 elif self.stop is not None:118 params["stop"] = self.stop119 else:120 params["stop"] = stop121 return {**params, **kwargs}122 123 def _stream(124 self,125 messages: List[BaseMessage],126 stop: Optional[List[str]] = None,127 run_manager: Optional[CallbackManagerForLLMRun] = None,128 **kwargs: Any,129 ) -> Iterator[ChatGenerationChunk]:130 params = self._get_invocation_params(stop=stop, **kwargs)131 stream = self.client.chat.completions.create(132 **get_chat_request(messages), stream=True, model=self.model, **params133 )134 for chunk in stream:135 delta = chunk.choices[0].delta.content136 if delta:137 if run_manager:138 run_manager.on_llm_new_token(delta)139 yield ChatGenerationChunk(message=AIMessageChunk(content=delta))140 141 async def _astream(142 self,143 messages: List[BaseMessage],144 stop: Optional[List[str]] = None,145 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,146 **kwargs: Any,147 ) -> AsyncIterator[ChatGenerationChunk]:148 params = self._get_invocation_params(stop=stop, **kwargs)149 stream = await self.async_client.chat.completions.create(150 **get_chat_request(messages), stream=True, model=self.model, **params151 )152 async for chunk in stream:153 delta = chunk.choices[0].delta.content154 if delta:155 if run_manager:156 await run_manager.on_llm_new_token(delta)157 yield ChatGenerationChunk(message=AIMessageChunk(content=delta))158 159 def _generate(160 self,161 messages: List[BaseMessage],162 stop: Optional[List[str]] = None,163 run_manager: Optional[CallbackManagerForLLMRun] = None,164 **kwargs: Any,165 ) -> ChatResult:166 if self.streaming:167 stream_iter = self._stream(168 messages, stop=stop, run_manager=run_manager, **kwargs169 )170 return generate_from_stream(stream_iter)171 172 params = self._get_invocation_params(stop=stop, **kwargs)173 response = self.client.chat.completions.create(174 messages=[175 {176 "role": get_role(message),177 "content": message.content,178 }179 for message in messages180 ],181 stream=False,182 model=self.model,183 **params,184 )185 186 message = AIMessage(content=response.choices[0].message.content)187 return ChatResult(generations=[ChatGeneration(message=message)])188 189 async def _agenerate(190 self,191 messages: List[BaseMessage],192 stop: Optional[List[str]] = None,193 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,194 **kwargs: Any,195 ) -> ChatResult:196 if self.streaming:197 stream_iter = self._astream(198 messages, stop=stop, run_manager=run_manager, **kwargs199 )200 return await agenerate_from_stream(stream_iter)201 202 params = self._get_invocation_params(stop=stop, **kwargs)203 response = await self.async_client.chat.completions.create(204 messages=[205 {206 "role": get_role(message),207 "content": message.content,208 }209 for message in messages210 ],211 stream=False,212 model=self.model,213 **params,214 )215 216 message = AIMessage(content=response.choices[0].message.content)217 return ChatResult(generations=[ChatGeneration(message=message)])218 