codekingpro/portable-devtools
114k
1import json2import warnings3from typing import (4 Any,5 AsyncIterator,6 Dict,7 Iterator,8 List,9 Mapping,10 Optional,11 Type,12 cast,13)14 15from langchain_core.callbacks import (16 AsyncCallbackManagerForLLMRun,17 CallbackManagerForLLMRun,18)19from langchain_core.language_models.chat_models import BaseChatModel20from langchain_core.messages import (21 AIMessage,22 AIMessageChunk,23 BaseMessage,24 BaseMessageChunk,25 ChatMessage,26 ChatMessageChunk,27 FunctionMessageChunk,28 HumanMessage,29 HumanMessageChunk,30 SystemMessage,31 SystemMessageChunk,32 ToolMessageChunk,33)34from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult35 36from langchain_community.llms.azureml_endpoint import (37 AzureMLBaseEndpoint,38 AzureMLEndpointApiType,39 ContentFormatterBase,40)41 42 43class LlamaContentFormatter(ContentFormatterBase):44 """Content formatter for `LLaMA`."""45 46 def __init__(self) -> None:47 raise TypeError(48 "`LlamaContentFormatter` is deprecated for chat models. Use "49 "`CustomOpenAIContentFormatter` instead."50 )51 52 53class CustomOpenAIChatContentFormatter(ContentFormatterBase):54 """Chat Content formatter for models with OpenAI like API scheme."""55 56 SUPPORTED_ROLES: List[str] = ["user", "assistant", "system"]57 58 @staticmethod59 def _convert_message_to_dict(message: BaseMessage) -> Dict:60 """Converts a message to a dict according to a role"""61 content = cast(str, message.content)62 if isinstance(message, HumanMessage):63 return {64 "role": "user",65 "content": ContentFormatterBase.escape_special_characters(content),66 }67 elif isinstance(message, AIMessage):68 return {69 "role": "assistant",70 "content": ContentFormatterBase.escape_special_characters(content),71 }72 elif isinstance(message, SystemMessage):73 return {74 "role": "system",75 "content": ContentFormatterBase.escape_special_characters(content),76 }77 elif (78 isinstance(message, ChatMessage)79 and message.role in CustomOpenAIChatContentFormatter.SUPPORTED_ROLES80 ):81 return {82 "role": message.role,83 "content": ContentFormatterBase.escape_special_characters(content),84 }85 else:86 supported = ",".join(87 [role for role in CustomOpenAIChatContentFormatter.SUPPORTED_ROLES]88 )89 raise ValueError(90 f"""Received unsupported role. 91 Supported roles for the LLaMa Foundation Model: {supported}"""92 )93 94 @property95 def supported_api_types(self) -> List[AzureMLEndpointApiType]:96 return [AzureMLEndpointApiType.dedicated, AzureMLEndpointApiType.serverless]97 98 def format_messages_request_payload(99 self,100 messages: List[BaseMessage],101 model_kwargs: Dict,102 api_type: AzureMLEndpointApiType,103 ) -> bytes:104 """Formats the request according to the chosen api"""105 chat_messages = [106 CustomOpenAIChatContentFormatter._convert_message_to_dict(message)107 for message in messages108 ]109 if api_type in [110 AzureMLEndpointApiType.dedicated,111 AzureMLEndpointApiType.realtime,112 ]:113 request_payload = json.dumps(114 {115 "input_data": {116 "input_string": chat_messages,117 "parameters": model_kwargs,118 }119 }120 )121 elif api_type == AzureMLEndpointApiType.serverless:122 request_payload = json.dumps({"messages": chat_messages, **model_kwargs})123 else:124 raise ValueError(125 f"`api_type` {api_type} is not supported by this formatter"126 )127 return str.encode(request_payload)128 129 def format_response_payload(130 self,131 output: bytes,132 api_type: AzureMLEndpointApiType = AzureMLEndpointApiType.dedicated,133 ) -> ChatGeneration:134 """Formats response"""135 if api_type in [136 AzureMLEndpointApiType.dedicated,137 AzureMLEndpointApiType.realtime,138 ]:139 try:140 choice = json.loads(output)["output"]141 except (KeyError, IndexError, TypeError) as e:142 raise ValueError(self.format_error_msg.format(api_type=api_type)) from e143 return ChatGeneration(144 message=AIMessage(145 content=choice.strip(),146 ),147 generation_info=None,148 )149 if api_type == AzureMLEndpointApiType.serverless:150 try:151 choice = json.loads(output)["choices"][0]152 if not isinstance(choice, dict):153 raise TypeError(154 "Endpoint response is not well formed for a chat "155 "model. Expected `dict` but `{type(choice)}` was received."156 )157 except (KeyError, IndexError, TypeError) as e:158 raise ValueError(self.format_error_msg.format(api_type=api_type)) from e159 return ChatGeneration(160 message=AIMessage(content=choice["message"]["content"].strip())161 if choice["message"]["role"] == "assistant"162 else BaseMessage(163 content=choice["message"]["content"].strip(),164 type=choice["message"]["role"],165 ),166 generation_info=dict(167 finish_reason=choice.get("finish_reason"),168 logprobs=choice.get("logprobs"),169 ),170 )171 raise ValueError(f"`api_type` {api_type} is not supported by this formatter")172 173 174class LlamaChatContentFormatter(CustomOpenAIChatContentFormatter):175 """Deprecated: Kept for backwards compatibility176 177 Chat Content formatter for Llama."""178 179 def __init__(self) -> None:180 super().__init__()181 warnings.warn(182 """`LlamaChatContentFormatter` will be deprecated in the future. 183 Please use `CustomOpenAIChatContentFormatter` instead. 184 """185 )186 187 188class MistralChatContentFormatter(LlamaChatContentFormatter):189 """Content formatter for `Mistral`."""190 191 def format_messages_request_payload(192 self,193 messages: List[BaseMessage],194 model_kwargs: Dict,195 api_type: AzureMLEndpointApiType,196 ) -> bytes:197 """Formats the request according to the chosen api"""198 chat_messages = [self._convert_message_to_dict(message) for message in messages]199 200 if chat_messages and chat_messages[0]["role"] == "system":201 # Mistral OSS models do not explicitly support system prompts, so we have to202 # stash in the first user prompt203 chat_messages[1]["content"] = (204 chat_messages[0]["content"] + "\n\n" + chat_messages[1]["content"]205 )206 del chat_messages[0]207 208 if api_type == AzureMLEndpointApiType.realtime:209 request_payload = json.dumps(210 {211 "input_data": {212 "input_string": chat_messages,213 "parameters": model_kwargs,214 }215 }216 )217 elif api_type == AzureMLEndpointApiType.serverless:218 request_payload = json.dumps({"messages": chat_messages, **model_kwargs})219 else:220 raise ValueError(221 f"`api_type` {api_type} is not supported by this formatter"222 )223 return str.encode(request_payload)224 225 226class AzureMLChatOnlineEndpoint(BaseChatModel, AzureMLBaseEndpoint):227 """Azure ML Online Endpoint chat models.228 229 Example:230 .. code-block:: python231 azure_llm = AzureMLOnlineEndpoint(232 endpoint_url="https://<your-endpoint>.<your_region>.inference.ml.azure.com/v1/chat/completions",233 endpoint_api_type=AzureMLApiType.serverless,234 endpoint_api_key="my-api-key",235 content_formatter=chat_content_formatter,236 )237 """238 239 @property240 def _identifying_params(self) -> Dict[str, Any]:241 """Get the identifying parameters."""242 _model_kwargs = self.model_kwargs or {}243 return {244 **{"model_kwargs": _model_kwargs},245 }246 247 @property248 def _llm_type(self) -> str:249 """Return type of llm."""250 return "azureml_chat_endpoint"251 252 def _generate(253 self,254 messages: List[BaseMessage],255 stop: Optional[List[str]] = None,256 run_manager: Optional[CallbackManagerForLLMRun] = None,257 **kwargs: Any,258 ) -> ChatResult:259 """Call out to an AzureML Managed Online endpoint.260 Args:261 messages: The messages in the conversation with the chat model.262 stop: Optional list of stop words to use when generating.263 Returns:264 The string generated by the model.265 Example:266 .. code-block:: python267 response = azureml_model.invoke("Tell me a joke.")268 """269 _model_kwargs = self.model_kwargs or {}270 _model_kwargs.update(kwargs)271 if stop:272 _model_kwargs["stop"] = stop273 274 request_payload = self.content_formatter.format_messages_request_payload(275 messages, _model_kwargs, self.endpoint_api_type276 )277 response_payload = self.http_client.call(278 body=request_payload, run_manager=run_manager279 )280 generations = self.content_formatter.format_response_payload(281 response_payload, self.endpoint_api_type282 )283 return ChatResult(generations=[generations])284 285 def _stream(286 self,287 messages: List[BaseMessage],288 stop: Optional[List[str]] = None,289 run_manager: Optional[CallbackManagerForLLMRun] = None,290 **kwargs: Any,291 ) -> Iterator[ChatGenerationChunk]:292 self.endpoint_url = self.endpoint_url.replace("/chat/completions", "")293 timeout = None if "timeout" not in kwargs else kwargs["timeout"]294 295 import openai296 297 params = {}298 client_params = {299 "api_key": self.endpoint_api_key.get_secret_value(),300 "base_url": self.endpoint_url,301 "timeout": timeout,302 "default_headers": None,303 "default_query": None,304 "http_client": None,305 }306 307 client = openai.OpenAI(**client_params)308 message_dicts = [309 CustomOpenAIChatContentFormatter._convert_message_to_dict(m)310 for m in messages311 ]312 params = {"stream": True, "stop": stop, "model": None, **kwargs}313 314 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk315 for chunk in client.chat.completions.create(messages=message_dicts, **params):316 if not isinstance(chunk, dict):317 chunk = chunk.dict()318 if len(chunk["choices"]) == 0:319 continue320 choice = chunk["choices"][0]321 chunk = _convert_delta_to_message_chunk(322 choice["delta"],323 default_chunk_class,324 )325 generation_info = {}326 if finish_reason := choice.get("finish_reason"):327 generation_info["finish_reason"] = finish_reason328 logprobs = choice.get("logprobs")329 if logprobs:330 generation_info["logprobs"] = logprobs331 default_chunk_class = chunk.__class__332 chunk = ChatGenerationChunk(333 message=chunk,334 generation_info=generation_info or None,335 )336 if run_manager:337 run_manager.on_llm_new_token(chunk.text, chunk=chunk, logprobs=logprobs)338 yield chunk339 340 async def _astream(341 self,342 messages: List[BaseMessage],343 stop: Optional[List[str]] = None,344 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,345 **kwargs: Any,346 ) -> AsyncIterator[ChatGenerationChunk]:347 self.endpoint_url = self.endpoint_url.replace("/chat/completions", "")348 timeout = None if "timeout" not in kwargs else kwargs["timeout"]349 350 import openai351 352 params = {}353 client_params = {354 "api_key": self.endpoint_api_key.get_secret_value(),355 "base_url": self.endpoint_url,356 "timeout": timeout,357 "default_headers": None,358 "default_query": None,359 "http_client": None,360 }361 362 async_client = openai.AsyncOpenAI(**client_params)363 message_dicts = [364 CustomOpenAIChatContentFormatter._convert_message_to_dict(m)365 for m in messages366 ]367 params = {"stream": True, "stop": stop, "model": None, **kwargs}368 369 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk370 async for chunk in await async_client.chat.completions.create(371 messages=message_dicts,372 **params,373 ):374 if not isinstance(chunk, dict):375 chunk = chunk.dict()376 if len(chunk["choices"]) == 0:377 continue378 choice = chunk["choices"][0]379 chunk = _convert_delta_to_message_chunk(380 choice["delta"], default_chunk_class381 )382 generation_info = {}383 if finish_reason := choice.get("finish_reason"):384 generation_info["finish_reason"] = finish_reason385 logprobs = choice.get("logprobs")386 if logprobs:387 generation_info["logprobs"] = logprobs388 default_chunk_class = chunk.__class__389 chunk = ChatGenerationChunk(390 message=chunk, generation_info=generation_info or None391 )392 if run_manager:393 await run_manager.on_llm_new_token(394 token=chunk.text, chunk=chunk, logprobs=logprobs395 )396 yield chunk397 398 399def _convert_delta_to_message_chunk(400 _dict: Mapping[str, Any], default_class: Type[BaseMessageChunk]401) -> BaseMessageChunk:402 role = cast(str, _dict.get("role"))403 content = cast(str, _dict.get("content") or "")404 additional_kwargs: Dict = {}405 if _dict.get("function_call"):406 function_call = dict(_dict["function_call"])407 if "name" in function_call and function_call["name"] is None:408 function_call["name"] = ""409 additional_kwargs["function_call"] = function_call410 if _dict.get("tool_calls"):411 additional_kwargs["tool_calls"] = _dict["tool_calls"]412 413 if role == "user" or default_class == HumanMessageChunk:414 return HumanMessageChunk(content=content)415 elif role == "assistant" or default_class == AIMessageChunk:416 return AIMessageChunk(content=content, additional_kwargs=additional_kwargs)417 elif role == "system" or default_class == SystemMessageChunk:418 return SystemMessageChunk(content=content)419 elif role == "function" or default_class == FunctionMessageChunk:420 return FunctionMessageChunk(content=content, name=_dict["name"])421 elif role == "tool" or default_class == ToolMessageChunk:422 return ToolMessageChunk(content=content, tool_call_id=_dict["tool_call_id"])423 elif role or default_class == ChatMessageChunk:424 return ChatMessageChunk(content=content, role=role)425 else:426 return default_class(content=content) # type: ignore[call-arg]427 