Team Ai
Datasetpublic

codekingpro/portable-devtools

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