Team Ai
Apppublic

aphilippov/python-server-api

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
multimodal_conversable_agent.py94 linesDownload Raw Back to contrib
1import copy2from typing import Any, Callable, Dict, List, Optional, Tuple, Union3 4from autogen import OpenAIWrapper5from autogen.agentchat import Agent, ConversableAgent6from autogen.agentchat.contrib.img_utils import gpt4v_formatter7 8try:9    from termcolor import colored10except ImportError:11 12    def colored(x, *args, **kwargs):13        return x14 15 16from autogen.code_utils import content_str17 18DEFAULT_LMM_SYS_MSG = """You are a helpful AI assistant."""19DEFAULT_MODEL = "gpt-4-vision-preview"20 21 22class MultimodalConversableAgent(ConversableAgent):23    DEFAULT_CONFIG = {24        "model": DEFAULT_MODEL,25    }26 27    def __init__(28        self,29        name: str,30        system_message: Optional[Union[str, List]] = DEFAULT_LMM_SYS_MSG,31        is_termination_msg: str = None,32        *args,33        **kwargs,34    ):35        """36        Args:37            name (str): agent name.38            system_message (str): system message for the OpenAIWrapper inference.39                Please override this attribute if you want to reprogram the agent.40            **kwargs (dict): Please refer to other kwargs in41                [ConversableAgent](../conversable_agent#__init__).42        """43        super().__init__(44            name,45            system_message,46            is_termination_msg=is_termination_msg,47            *args,48            **kwargs,49        )50        # call the setter to handle special format.51        self.update_system_message(system_message)52        self._is_termination_msg = (53            is_termination_msg54            if is_termination_msg is not None55            else (lambda x: content_str(x.get("content")) == "TERMINATE")56        )57 58    def update_system_message(self, system_message: Union[Dict, List, str]):59        """Update the system message.60 61        Args:62            system_message (str): system message for the OpenAIWrapper inference.63        """64        self._oai_system_message[0]["content"] = self._message_to_dict(system_message)["content"]65        self._oai_system_message[0]["role"] = "system"66 67    @staticmethod68    def _message_to_dict(message: Union[Dict, List, str]) -> Dict:69        """Convert a message to a dictionary. This implementation70        handles the GPT-4V formatting for easier prompts.71 72        The message can be a string, a dictionary, or a list of dictionaries:73            - If it's a string, it will be cast into a list and placed in the 'content' field.74            - If it's a list, it will be directly placed in the 'content' field.75            - If it's a dictionary, it is already in message dict format. The 'content' field of this dictionary76            will be processed using the gpt4v_formatter.77        """78        if isinstance(message, str):79            return {"content": gpt4v_formatter(message)}80        if isinstance(message, list):81            return {"content": message}82        if isinstance(message, dict):83            assert "content" in message, "The message dict must have a `content` field"84            if isinstance(message["content"], str):85                message = copy.deepcopy(message)86                message["content"] = gpt4v_formatter(message["content"])87            try:88                content_str(message["content"])89            except (TypeError, ValueError) as e:90                print("The `content` field should be compatible with the content_str function!")91                raise e92            return message93        raise ValueError(f"Unsupported message type: {type(message)}")94