Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
bedrock.py338 linesDownload Raw Back to chat_models
1import re2from collections import defaultdict3from typing import Any, Dict, Iterator, List, Optional, Tuple, Union4 5from langchain_core._api.deprecation import deprecated6from langchain_core.callbacks import (7    CallbackManagerForLLMRun,8)9from langchain_core.language_models.chat_models import BaseChatModel10from langchain_core.messages import (11    AIMessage,12    AIMessageChunk,13    BaseMessage,14    ChatMessage,15    HumanMessage,16    SystemMessage,17)18from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult19from pydantic import ConfigDict20 21from langchain_community.chat_models.anthropic import (22    convert_messages_to_prompt_anthropic,23)24from langchain_community.chat_models.meta import convert_messages_to_prompt_llama25from langchain_community.llms.bedrock import BedrockBase26from langchain_community.utilities.anthropic import (27    get_num_tokens_anthropic,28    get_token_ids_anthropic,29)30 31 32def _convert_one_message_to_text_mistral(message: BaseMessage) -> str:33    if isinstance(message, ChatMessage):34        message_text = f"\n\n{message.role.capitalize()}: {message.content}"35    elif isinstance(message, HumanMessage):36        message_text = f"[INST] {message.content} [/INST]"37    elif isinstance(message, AIMessage):38        message_text = f"{message.content}"39    elif isinstance(message, SystemMessage):40        message_text = f"<<SYS>> {message.content} <</SYS>>"41    else:42        raise ValueError(f"Got unknown type {message}")43    return message_text44 45 46def convert_messages_to_prompt_mistral(messages: List[BaseMessage]) -> str:47    """Convert a list of messages to a prompt for mistral."""48    return "\n".join(49        [_convert_one_message_to_text_mistral(message) for message in messages]50    )51 52 53def _format_image(image_url: str) -> Dict:54    """55    Formats an image of format data:image/jpeg;base64,{b64_string}56    to a dict for anthropic api57 58    {59      "type": "base64",60      "media_type": "image/jpeg",61      "data": "/9j/4AAQSkZJRg...",62    }63 64    And throws an error if it's not a b64 image65    """66    regex = r"^data:(?P<media_type>image/.+);base64,(?P<data>.+)$"67    match = re.match(regex, image_url)68    if match is None:69        raise ValueError(70            "Anthropic only supports base64-encoded images currently."71            " Example: data:image/png;base64,'/9j/4AAQSk'..."72        )73    return {74        "type": "base64",75        "media_type": match.group("media_type"),76        "data": match.group("data"),77    }78 79 80def _format_anthropic_messages(81    messages: List[BaseMessage],82) -> Tuple[Optional[str], List[Dict]]:83    """Format messages for anthropic."""84 85    """86    [87        {88            "role": _message_type_lookups[m.type],89            "content": [_AnthropicMessageContent(text=m.content).dict()],90        }91        for m in messages92    ]93    """94    system: Optional[str] = None95    formatted_messages: List[Dict] = []96    for i, message in enumerate(messages):97        if message.type == "system":98            if i != 0:99                raise ValueError("System message must be at beginning of message list.")100            if not isinstance(message.content, str):101                raise ValueError(102                    "System message must be a string, "103                    f"instead was: {type(message.content)}"104                )105            system = message.content106            continue107 108        role = _message_type_lookups[message.type]109        content: Union[str, List[Dict]]110 111        if not isinstance(message.content, str):112            # parse as dict113            assert isinstance(message.content, list), (114                "Anthropic message content must be str or list of dicts"115            )116 117            # populate content118            content = []119            for item in message.content:120                if isinstance(item, str):121                    content.append(122                        {123                            "type": "text",124                            "text": item,125                        }126                    )127                elif isinstance(item, dict):128                    if "type" not in item:129                        raise ValueError("Dict content item must have a type key")130                    if item["type"] == "image_url":131                        # convert format132                        source = _format_image(item["image_url"]["url"])133                        content.append(134                            {135                                "type": "image",136                                "source": source,137                            }138                        )139                    else:140                        content.append(item)141                else:142                    raise ValueError(143                        f"Content items must be str or dict, instead was: {type(item)}"144                    )145        else:146            content = message.content147 148        formatted_messages.append(149            {150                "role": role,151                "content": content,152            }153        )154    return system, formatted_messages155 156 157class ChatPromptAdapter:158    """Adapter class to prepare the inputs from Langchain to prompt format159    that Chat model expects.160    """161 162    @classmethod163    def convert_messages_to_prompt(164        cls, provider: str, messages: List[BaseMessage]165    ) -> str:166        if provider == "anthropic":167            prompt = convert_messages_to_prompt_anthropic(messages=messages)168        elif provider == "meta":169            prompt = convert_messages_to_prompt_llama(messages=messages)170        elif provider == "mistral":171            prompt = convert_messages_to_prompt_mistral(messages=messages)172        elif provider == "amazon":173            prompt = convert_messages_to_prompt_anthropic(174                messages=messages,175                human_prompt="\n\nUser:",176                ai_prompt="\n\nBot:",177            )178        else:179            raise NotImplementedError(180                f"Provider {provider} model does not support chat."181            )182        return prompt183 184    @classmethod185    def format_messages(186        cls, provider: str, messages: List[BaseMessage]187    ) -> Tuple[Optional[str], List[Dict]]:188        if provider == "anthropic":189            return _format_anthropic_messages(messages)190 191        raise NotImplementedError(192            f"Provider {provider} not supported for format_messages"193        )194 195 196_message_type_lookups = {197    "human": "user",198    "ai": "assistant",199    "AIMessageChunk": "assistant",200    "HumanMessageChunk": "user",201    "function": "user",202}203 204 205@deprecated(206    since="0.0.34", removal="1.0", alternative_import="langchain_aws.ChatBedrock"207)208class BedrockChat(BaseChatModel, BedrockBase):209    """Chat model that uses the Bedrock API."""210 211    @property212    def _llm_type(self) -> str:213        """Return type of chat model."""214        return "amazon_bedrock_chat"215 216    @classmethod217    def is_lc_serializable(cls) -> bool:218        """Return whether this model can be serialized by Langchain."""219        return True220 221    @classmethod222    def get_lc_namespace(cls) -> List[str]:223        """Get the namespace of the langchain object."""224        return ["langchain", "chat_models", "bedrock"]225 226    @property227    def lc_attributes(self) -> Dict[str, Any]:228        attributes: Dict[str, Any] = {}229 230        if self.region_name:231            attributes["region_name"] = self.region_name232 233        return attributes234 235    model_config = ConfigDict(236        extra="forbid",237    )238 239    def _stream(240        self,241        messages: List[BaseMessage],242        stop: Optional[List[str]] = None,243        run_manager: Optional[CallbackManagerForLLMRun] = None,244        **kwargs: Any,245    ) -> Iterator[ChatGenerationChunk]:246        provider = self._get_provider()247        prompt, system, formatted_messages = None, None, None248 249        if provider == "anthropic":250            system, formatted_messages = ChatPromptAdapter.format_messages(251                provider, messages252            )253        else:254            prompt = ChatPromptAdapter.convert_messages_to_prompt(255                provider=provider, messages=messages256            )257 258        for chunk in self._prepare_input_and_invoke_stream(259            prompt=prompt,260            system=system,261            messages=formatted_messages,262            stop=stop,263            run_manager=run_manager,264            **kwargs,265        ):266            delta = chunk.text267            yield ChatGenerationChunk(message=AIMessageChunk(content=delta))268 269    def _generate(270        self,271        messages: List[BaseMessage],272        stop: Optional[List[str]] = None,273        run_manager: Optional[CallbackManagerForLLMRun] = None,274        **kwargs: Any,275    ) -> ChatResult:276        completion = ""277        llm_output: Dict[str, Any] = {"model_id": self.model_id}278 279        if self.streaming:280            for chunk in self._stream(messages, stop, run_manager, **kwargs):281                completion += chunk.text282        else:283            provider = self._get_provider()284            prompt, system, formatted_messages = None, None, None285            params: Dict[str, Any] = {**kwargs}286 287            if provider == "anthropic":288                system, formatted_messages = ChatPromptAdapter.format_messages(289                    provider, messages290                )291            else:292                prompt = ChatPromptAdapter.convert_messages_to_prompt(293                    provider=provider, messages=messages294                )295 296            if stop:297                params["stop_sequences"] = stop298 299            completion, usage_info = self._prepare_input_and_invoke(300                prompt=prompt,301                stop=stop,302                run_manager=run_manager,303                system=system,304                messages=formatted_messages,305                **params,306            )307 308            llm_output["usage"] = usage_info309 310        return ChatResult(311            generations=[ChatGeneration(message=AIMessage(content=completion))],312            llm_output=llm_output,313        )314 315    def _combine_llm_outputs(self, llm_outputs: List[Optional[dict]]) -> dict:316        final_usage: Dict[str, int] = defaultdict(int)317        final_output = {}318        for output in llm_outputs:319            output = output or {}320            usage = output.get("usage", {})321            for token_type, token_count in usage.items():322                final_usage[token_type] += token_count323            final_output.update(output)324        final_output["usage"] = final_usage325        return final_output326 327    def get_num_tokens(self, text: str) -> int:328        if self._model_is_anthropic:329            return get_num_tokens_anthropic(text)330        else:331            return super().get_num_tokens(text)332 333    def get_token_ids(self, text: str) -> List[int]:334        if self._model_is_anthropic:335            return get_token_ids_anthropic(text)336        else:337            return super().get_token_ids(text)338 
codekingpro/portable-devtools · Team Ai