Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model_invocation_utils.py171 linesDownload Raw Back to utils
1"""2For some reason, model will be used in tools like WebScraperTool, WikipediaSearchTool etc.3 4Therefore, a model manager is needed to list/invoke/validate models.5"""6 7import json8from typing import cast9 10from core.model_manager import ModelManager11from core.model_runtime.entities.llm_entities import LLMResult12from core.model_runtime.entities.message_entities import PromptMessage13from core.model_runtime.entities.model_entities import ModelType14from core.model_runtime.errors.invoke import (15    InvokeAuthorizationError,16    InvokeBadRequestError,17    InvokeConnectionError,18    InvokeRateLimitError,19    InvokeServerUnavailableError,20)21from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel, ModelPropertyKey22from core.model_runtime.utils.encoders import jsonable_encoder23from extensions.ext_database import db24from models.tools import ToolModelInvoke25 26 27class InvokeModelError(Exception):28    pass29 30 31class ModelInvocationUtils:32    @staticmethod33    def get_max_llm_context_tokens(34        tenant_id: str,35    ) -> int:36        """37        get max llm context tokens of the model38        """39        model_manager = ModelManager()40        model_instance = model_manager.get_default_model_instance(41            tenant_id=tenant_id,42            model_type=ModelType.LLM,43        )44 45        if not model_instance:46            raise InvokeModelError("Model not found")47 48        llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)49        schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)50 51        if not schema:52            raise InvokeModelError("No model schema found")53 54        max_tokens = schema.model_properties.get(ModelPropertyKey.CONTEXT_SIZE, None)55        if max_tokens is None:56            return 204857 58        return max_tokens59 60    @staticmethod61    def calculate_tokens(tenant_id: str, prompt_messages: list[PromptMessage]) -> int:62        """63        calculate tokens from prompt messages and model parameters64        """65 66        # get model instance67        model_manager = ModelManager()68        model_instance = model_manager.get_default_model_instance(tenant_id=tenant_id, model_type=ModelType.LLM)69 70        if not model_instance:71            raise InvokeModelError("Model not found")72 73        # get tokens74        tokens = model_instance.get_llm_num_tokens(prompt_messages)75 76        return tokens77 78    @staticmethod79    def invoke(80        user_id: str, tenant_id: str, tool_type: str, tool_name: str, prompt_messages: list[PromptMessage]81    ) -> LLMResult:82        """83        invoke model with parameters in user's own context84 85        :param user_id: user id86        :param tenant_id: tenant id, the tenant id of the creator of the tool87        :param tool_provider: tool provider88        :param tool_id: tool id89        :param tool_name: tool name90        :param provider: model provider91        :param model: model name92        :param model_parameters: model parameters93        :param prompt_messages: prompt messages94        :return: AssistantPromptMessage95        """96 97        # get model manager98        model_manager = ModelManager()99        # get model instance100        model_instance = model_manager.get_default_model_instance(101            tenant_id=tenant_id,102            model_type=ModelType.LLM,103        )104 105        # get prompt tokens106        prompt_tokens = model_instance.get_llm_num_tokens(prompt_messages)107 108        model_parameters = {109            "temperature": 0.8,110            "top_p": 0.8,111        }112 113        # create tool model invoke114        tool_model_invoke = ToolModelInvoke(115            user_id=user_id,116            tenant_id=tenant_id,117            provider=model_instance.provider,118            tool_type=tool_type,119            tool_name=tool_name,120            model_parameters=json.dumps(model_parameters),121            prompt_messages=json.dumps(jsonable_encoder(prompt_messages)),122            model_response="",123            prompt_tokens=prompt_tokens,124            answer_tokens=0,125            answer_unit_price=0,126            answer_price_unit=0,127            provider_response_latency=0,128            total_price=0,129            currency="USD",130        )131 132        db.session.add(tool_model_invoke)133        db.session.commit()134 135        try:136            response: LLMResult = model_instance.invoke_llm(137                prompt_messages=prompt_messages,138                model_parameters=model_parameters,139                tools=[],140                stop=[],141                stream=False,142                user=user_id,143                callbacks=[],144            )145        except InvokeRateLimitError as e:146            raise InvokeModelError(f"Invoke rate limit error: {e}")147        except InvokeBadRequestError as e:148            raise InvokeModelError(f"Invoke bad request error: {e}")149        except InvokeConnectionError as e:150            raise InvokeModelError(f"Invoke connection error: {e}")151        except InvokeAuthorizationError as e:152            raise InvokeModelError("Invoke authorization error")153        except InvokeServerUnavailableError as e:154            raise InvokeModelError(f"Invoke server unavailable error: {e}")155        except Exception as e:156            raise InvokeModelError(f"Invoke error: {e}")157 158        # update tool model invoke159        tool_model_invoke.model_response = response.message.content160        if response.usage:161            tool_model_invoke.answer_tokens = response.usage.completion_tokens162            tool_model_invoke.answer_unit_price = response.usage.completion_unit_price163            tool_model_invoke.answer_price_unit = response.usage.completion_price_unit164            tool_model_invoke.provider_response_latency = response.usage.latency165            tool_model_invoke.total_price = response.usage.total_price166            tool_model_invoke.currency = response.usage.currency167 168        db.session.commit()169 170        return response171