Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
simple_prompt_transform.py327 linesDownload Raw Back to prompt
1import enum2import json3import os4from typing import TYPE_CHECKING, Optional5 6from core.app.app_config.entities import PromptTemplateEntity7from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity8from core.file import file_manager9from core.memory.token_buffer_memory import TokenBufferMemory10from core.model_runtime.entities.message_entities import (11    PromptMessage,12    PromptMessageContent,13    SystemPromptMessage,14    TextPromptMessageContent,15    UserPromptMessage,16)17from core.prompt.entities.advanced_prompt_entities import MemoryConfig18from core.prompt.prompt_transform import PromptTransform19from core.prompt.utils.prompt_template_parser import PromptTemplateParser20from models.model import AppMode21 22if TYPE_CHECKING:23    from core.file.models import File24 25 26class ModelMode(str, enum.Enum):27    COMPLETION = "completion"28    CHAT = "chat"29 30    @classmethod31    def value_of(cls, value: str) -> "ModelMode":32        """33        Get value of given mode.34 35        :param value: mode value36        :return: mode37        """38        for mode in cls:39            if mode.value == value:40                return mode41        raise ValueError(f"invalid mode value {value}")42 43 44prompt_file_contents = {}45 46 47class SimplePromptTransform(PromptTransform):48    """49    Simple Prompt Transform for Chatbot App Basic Mode.50    """51 52    def get_prompt(53        self,54        app_mode: AppMode,55        prompt_template_entity: PromptTemplateEntity,56        inputs: dict,57        query: str,58        files: list["File"],59        context: Optional[str],60        memory: Optional[TokenBufferMemory],61        model_config: ModelConfigWithCredentialsEntity,62    ) -> tuple[list[PromptMessage], Optional[list[str]]]:63        inputs = {key: str(value) for key, value in inputs.items()}64 65        model_mode = ModelMode.value_of(model_config.mode)66        if model_mode == ModelMode.CHAT:67            prompt_messages, stops = self._get_chat_model_prompt_messages(68                app_mode=app_mode,69                pre_prompt=prompt_template_entity.simple_prompt_template,70                inputs=inputs,71                query=query,72                files=files,73                context=context,74                memory=memory,75                model_config=model_config,76            )77        else:78            prompt_messages, stops = self._get_completion_model_prompt_messages(79                app_mode=app_mode,80                pre_prompt=prompt_template_entity.simple_prompt_template,81                inputs=inputs,82                query=query,83                files=files,84                context=context,85                memory=memory,86                model_config=model_config,87            )88 89        return prompt_messages, stops90 91    def get_prompt_str_and_rules(92        self,93        app_mode: AppMode,94        model_config: ModelConfigWithCredentialsEntity,95        pre_prompt: str,96        inputs: dict,97        query: Optional[str] = None,98        context: Optional[str] = None,99        histories: Optional[str] = None,100    ) -> tuple[str, dict]:101        # get prompt template102        prompt_template_config = self.get_prompt_template(103            app_mode=app_mode,104            provider=model_config.provider,105            model=model_config.model,106            pre_prompt=pre_prompt,107            has_context=context is not None,108            query_in_prompt=query is not None,109            with_memory_prompt=histories is not None,110        )111 112        variables = {k: inputs[k] for k in prompt_template_config["custom_variable_keys"] if k in inputs}113 114        for v in prompt_template_config["special_variable_keys"]:115            # support #context#, #query# and #histories#116            if v == "#context#":117                variables["#context#"] = context or ""118            elif v == "#query#":119                variables["#query#"] = query or ""120            elif v == "#histories#":121                variables["#histories#"] = histories or ""122 123        prompt_template = prompt_template_config["prompt_template"]124        prompt = prompt_template.format(variables)125 126        return prompt, prompt_template_config["prompt_rules"]127 128    def get_prompt_template(129        self,130        app_mode: AppMode,131        provider: str,132        model: str,133        pre_prompt: str,134        has_context: bool,135        query_in_prompt: bool,136        with_memory_prompt: bool = False,137    ) -> dict:138        prompt_rules = self._get_prompt_rule(app_mode=app_mode, provider=provider, model=model)139 140        custom_variable_keys = []141        special_variable_keys = []142 143        prompt = ""144        for order in prompt_rules["system_prompt_orders"]:145            if order == "context_prompt" and has_context:146                prompt += prompt_rules["context_prompt"]147                special_variable_keys.append("#context#")148            elif order == "pre_prompt" and pre_prompt:149                prompt += pre_prompt + "\n"150                pre_prompt_template = PromptTemplateParser(template=pre_prompt)151                custom_variable_keys = pre_prompt_template.variable_keys152            elif order == "histories_prompt" and with_memory_prompt:153                prompt += prompt_rules["histories_prompt"]154                special_variable_keys.append("#histories#")155 156        if query_in_prompt:157            prompt += prompt_rules.get("query_prompt", "{{#query#}}")158            special_variable_keys.append("#query#")159 160        return {161            "prompt_template": PromptTemplateParser(template=prompt),162            "custom_variable_keys": custom_variable_keys,163            "special_variable_keys": special_variable_keys,164            "prompt_rules": prompt_rules,165        }166 167    def _get_chat_model_prompt_messages(168        self,169        app_mode: AppMode,170        pre_prompt: str,171        inputs: dict,172        query: str,173        context: Optional[str],174        files: list["File"],175        memory: Optional[TokenBufferMemory],176        model_config: ModelConfigWithCredentialsEntity,177    ) -> tuple[list[PromptMessage], Optional[list[str]]]:178        prompt_messages = []179 180        # get prompt181        prompt, _ = self.get_prompt_str_and_rules(182            app_mode=app_mode,183            model_config=model_config,184            pre_prompt=pre_prompt,185            inputs=inputs,186            query=None,187            context=context,188        )189 190        if prompt and query:191            prompt_messages.append(SystemPromptMessage(content=prompt))192 193        if memory:194            prompt_messages = self._append_chat_histories(195                memory=memory,196                memory_config=MemoryConfig(197                    window=MemoryConfig.WindowConfig(198                        enabled=False,199                    )200                ),201                prompt_messages=prompt_messages,202                model_config=model_config,203            )204 205        if query:206            prompt_messages.append(self.get_last_user_message(query, files))207        else:208            prompt_messages.append(self.get_last_user_message(prompt, files))209 210        return prompt_messages, None211 212    def _get_completion_model_prompt_messages(213        self,214        app_mode: AppMode,215        pre_prompt: str,216        inputs: dict,217        query: str,218        context: Optional[str],219        files: list["File"],220        memory: Optional[TokenBufferMemory],221        model_config: ModelConfigWithCredentialsEntity,222    ) -> tuple[list[PromptMessage], Optional[list[str]]]:223        # get prompt224        prompt, prompt_rules = self.get_prompt_str_and_rules(225            app_mode=app_mode,226            model_config=model_config,227            pre_prompt=pre_prompt,228            inputs=inputs,229            query=query,230            context=context,231        )232 233        if memory:234            tmp_human_message = UserPromptMessage(content=prompt)235 236            rest_tokens = self._calculate_rest_token([tmp_human_message], model_config)237            histories = self._get_history_messages_from_memory(238                memory=memory,239                memory_config=MemoryConfig(240                    window=MemoryConfig.WindowConfig(241                        enabled=False,242                    )243                ),244                max_token_limit=rest_tokens,245                human_prefix=prompt_rules.get("human_prefix", "Human"),246                ai_prefix=prompt_rules.get("assistant_prefix", "Assistant"),247            )248 249            # get prompt250            prompt, prompt_rules = self.get_prompt_str_and_rules(251                app_mode=app_mode,252                model_config=model_config,253                pre_prompt=pre_prompt,254                inputs=inputs,255                query=query,256                context=context,257                histories=histories,258            )259 260        stops = prompt_rules.get("stops")261        if stops is not None and len(stops) == 0:262            stops = None263 264        return [self.get_last_user_message(prompt, files)], stops265 266    def get_last_user_message(self, prompt: str, files: list["File"]) -> UserPromptMessage:267        if files:268            prompt_message_contents: list[PromptMessageContent] = []269            prompt_message_contents.append(TextPromptMessageContent(data=prompt))270            for file in files:271                prompt_message_contents.append(file_manager.to_prompt_message_content(file))272 273            prompt_message = UserPromptMessage(content=prompt_message_contents)274        else:275            prompt_message = UserPromptMessage(content=prompt)276 277        return prompt_message278 279    def _get_prompt_rule(self, app_mode: AppMode, provider: str, model: str) -> dict:280        """281        Get simple prompt rule.282        :param app_mode: app mode283        :param provider: model provider284        :param model: model name285        :return:286        """287        prompt_file_name = self._prompt_file_name(app_mode=app_mode, provider=provider, model=model)288 289        # Check if the prompt file is already loaded290        if prompt_file_name in prompt_file_contents:291            return prompt_file_contents[prompt_file_name]292 293        # Get the absolute path of the subdirectory294        prompt_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "prompt_templates")295        json_file_path = os.path.join(prompt_path, f"{prompt_file_name}.json")296 297        # Open the JSON file and read its content298        with open(json_file_path, encoding="utf-8") as json_file:299            content = json.load(json_file)300 301            # Store the content of the prompt file302            prompt_file_contents[prompt_file_name] = content303 304            return content305 306    def _prompt_file_name(self, app_mode: AppMode, provider: str, model: str) -> str:307        # baichuan308        is_baichuan = False309        if provider == "baichuan":310            is_baichuan = True311        else:312            baichuan_supported_providers = ["huggingface_hub", "openllm", "xinference"]313            if provider in baichuan_supported_providers and "baichuan" in model.lower():314                is_baichuan = True315 316        if is_baichuan:317            if app_mode == AppMode.COMPLETION:318                return "baichuan_completion"319            else:320                return "baichuan_chat"321 322        # common323        if app_mode == AppMode.COMPLETION:324            return "common_completion"325        else:326            return "common_chat"327