diff --git a/apps/models_provider/impl/qianfan_model_provider/__init__.py b/apps/models_provider/impl/qianfan_model_provider/__init__.py new file mode 100644 index 00000000000..fd54226fe4c --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/__init__.py @@ -0,0 +1,8 @@ +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: __init__.py.py +@date:2023/10/31 17:16 +@desc: +""" diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/__init__.py b/apps/models_provider/impl/qianfan_model_provider/credential/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/embedding.py b/apps/models_provider/impl/qianfan_model_provider/credential/embedding.py new file mode 100644 index 00000000000..d5e1c7b9e0c --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/embedding.py @@ -0,0 +1,58 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author:虎 +@file: embedding.py +@date:2024/10/17 15:40 +@desc: +""" + +from typing import Dict + +from django.utils.translation import gettext as _ + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class QianfanEmbeddingCredential(BaseForm, BaseModelCredential): + api_base = forms.TextInputField("API URL", required=True, default_value="https://qianfan.baidubce.com/v2") + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, object], + model_params, + provider, + raise_exception=False, + ): + for key in ["api_base", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, _("{key} is required").format(key=key)) + return False + + try: + model = provider.get_model(model_type, model_name, model_credential) + model.embed_query(_("Hello")) + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + if isinstance(e, AppApiException): + raise e + if raise_exception: + raise AppApiException( + ValidCode.valid_error.value, + _("Verification failed, please check whether the parameters are correct: {error}").format( + error=str(e) + ), + ) + return False + return True + + def encryption_dict(self, model: Dict[str, object]): + return {**model, "api_key": super().encryption(model.get("api_key", ""))} diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/image.py b/apps/models_provider/impl/qianfan_model_provider/credential/image.py new file mode 100644 index 00000000000..02faa0d5847 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/image.py @@ -0,0 +1,89 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: image.py +@desc: 千帆视觉理解模型凭据 +""" + +from typing import Dict + +from django.utils.translation import gettext, gettext_lazy as _ +from langchain_core.messages import HumanMessage + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, TooltipLabel +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class QianfanImageModelParams(BaseForm): + temperature = forms.SliderField( + TooltipLabel( + _("Temperature"), + _("Higher values make the output more random, while lower values make it more focused and deterministic"), + ), + required=True, + default_value=0.95, + _min=0.1, + _max=1.0, + _step=0.01, + precision=2, + ) + + max_tokens = forms.SliderField( + TooltipLabel( + _("Output the maximum Tokens"), _("Specify the maximum number of tokens that the model can generate") + ), + required=True, + default_value=1024, + _min=1, + _max=100000, + _step=1, + precision=0, + ) + + +class QianfanImageModelCredential(BaseForm, BaseModelCredential): + api_base = forms.TextInputField("API URL", required=True, default_value="https://qianfan.baidubce.com/v2") + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, object], + model_params, + provider, + raise_exception=False, + ): + for key in ["api_base", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key=key)) + return False + + try: + model = provider.get_model(model_type, model_name, model_credential, **model_params) + response = model.stream([HumanMessage(content=[{"type": "text", "text": gettext("Hello")}])]) + for _chunk in response: + break + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + if isinstance(e, AppApiException): + raise e + if raise_exception: + raise AppApiException( + ValidCode.valid_error.value, + gettext("Verification failed, please check whether the parameters are correct: {error}").format( + error=str(e) + ), + ) + return False + return True + + def encryption_dict(self, model: Dict[str, object]): + return {**model, "api_key": super().encryption(model.get("api_key", ""))} + + def get_model_params_setting_form(self, model_name): + return QianfanImageModelParams() diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/llm.py b/apps/models_provider/impl/qianfan_model_provider/credential/llm.py new file mode 100644 index 00000000000..5cc2afb467b --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/llm.py @@ -0,0 +1,88 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author:虎 +@file: llm.py +@date:2024/7/12 10:19 +@desc: +""" + +from typing import Dict + +from django.utils.translation import gettext, gettext_lazy as _ +from langchain_core.messages import HumanMessage + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, TooltipLabel +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class QianfanLLMModelParams(BaseForm): + temperature = forms.SliderField( + TooltipLabel( + _("Temperature"), + _("Higher values make the output more random, while lower values make it more focused and deterministic"), + ), + required=True, + default_value=0.95, + _min=0.1, + _max=1.0, + _step=0.01, + precision=2, + ) + + max_tokens = forms.SliderField( + TooltipLabel( + _("Output the maximum Tokens"), _("Specify the maximum number of tokens that the model can generate") + ), + required=True, + default_value=1024, + _min=2, + _max=100000, + _step=1, + precision=0, + ) + + +class QianfanLLMModelCredential(BaseForm, BaseModelCredential): + api_base = forms.TextInputField("API URL", required=True, default_value="https://qianfan.baidubce.com/v2") + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, object], + model_params, + provider, + raise_exception=False, + ): + for key in ["api_base", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key=key)) + return False + + try: + model = provider.get_model(model_type, model_name, model_credential, **{**model_params, "max_tokens": 1}) + model.invoke([HumanMessage(content="1")]) + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + raise e + return True + + def encryption_dict(self, model_info: Dict[str, object]): + return {**model_info, "api_key": super().encryption(model_info.get("api_key", ""))} + + def build_model(self, model_info: Dict[str, object]): + for key in ["api_base", "api_key", "model"]: + if key not in model_info: + raise AppApiException(500, gettext("{key} is required").format(key=key)) + self.api_base = model_info.get("api_base") + self.api_key = model_info.get("api_key") + return self + + def get_model_params_setting_form(self, model_name): + return QianfanLLMModelParams() diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/reranker.py b/apps/models_provider/impl/qianfan_model_provider/credential/reranker.py new file mode 100644 index 00000000000..b1f81512388 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/reranker.py @@ -0,0 +1,74 @@ +from typing import Dict + +from langchain_core.documents import Document + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, TooltipLabel +from models_provider.base_model_provider import BaseModelCredential, ValidCode +from django.utils.translation import gettext_lazy as _ +from common.utils.logger import maxkb_logger +from models_provider.impl.qianfan_model_provider.model.reranker import QfBgeReranker + + +class QfRerankerModelParams(BaseForm): + top_n = forms.SliderField( + TooltipLabel(_("Top N"), _("Number of top documents to return after reranking")), + required=True, + default_value=3, + _min=1, + _max=100, + _step=1, + precision=0, + ) + + +class QfRerankerCredential(BaseForm, BaseModelCredential): + api_url = forms.TextInputField("API URL", required=True) + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, object], + model_params, + provider, + raise_exception=True, + ): + model_type_list = provider.get_model_type_list() + if not any(list(filter(lambda mt: mt.get("value") == model_type, model_type_list))): + raise AppApiException( + ValidCode.valid_error.value, _("{model_type} Model type is not supported").format(model_type=model_type) + ) + + for key in ["api_url", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, _("{key} is required").format(key=key)) + else: + return False + try: + model: QfBgeReranker = provider.get_model(model_type, model_name, model_credential) + test_text = str(_("Hello")) + model.compress_documents([Document(page_content=test_text)], test_text) + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + if isinstance(e, AppApiException): + raise e + if raise_exception: + raise AppApiException( + ValidCode.valid_error.value, + _("Verification failed, please check whether the parameters are correct: {error}").format( + error=str(e) + ), + ) + return False + + return True + + def encryption_dict(self, model_info: Dict[str, object]): + return {**model_info, "api_key": super().encryption(model_info.get("api_key", ""))} + + def get_model_params_setting_form(self, model_name: str) -> QfRerankerModelParams: + return QfRerankerModelParams() diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/tti.py b/apps/models_provider/impl/qianfan_model_provider/credential/tti.py new file mode 100644 index 00000000000..27956275215 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/tti.py @@ -0,0 +1,106 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: tti.py +@desc: 千帆文生图模型凭据 +""" + +from typing import Dict + +from django.utils.translation import gettext, gettext_lazy as _ + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, TooltipLabel +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class QianfanTTIModelParams(BaseForm): + size = forms.SingleSelect( + TooltipLabel( + _("Image size"), + _( + "The size of the generated image. Optional values: [1024x1024, 1280x720, 720x1280, 1152x864, 864x1152, " + "1328x1328, 1664x928, 928x1664, 1472x1104, 1104x1472], default is 1024x1024." + ), + ), + required=True, + default_value="1024x1024", + option_list=[ + {"value": "1024x1024", "label": "1024x1024"}, + {"value": "1280x720", "label": "1280x720"}, + {"value": "720x1280", "label": "720x1280"}, + {"value": "1152x864", "label": "1152x864"}, + {"value": "864x1152", "label": "864x1152"}, + {"value": "1328x1328", "label": "1328x1328"}, + {"value": "1664x928", "label": "1664x928"}, + {"value": "928x1664", "label": "928x1664"}, + {"value": "1472x1104", "label": "1472x1104"}, + {"value": "1104x1472", "label": "1104x1472"}, + ], + text_field="label", + value_field="value", + ) + + response_format = forms.SingleSelect( + TooltipLabel( + _("Response format"), + _("The format of the generated image. url returns a URL, b64_json returns base64-encoded data."), + ), + required=True, + default_value="url", + option_list=[ + {"value": "url", "label": "url"}, + {"value": "b64_json", "label": "b64_json"}, + ], + text_field="label", + value_field="value", + ) + + +class QianfanTextToImageModelCredential(BaseForm, BaseModelCredential): + api_base = forms.TextInputField( + "API URL", + required=True, + default_value="https://qianfan.baidubce.com/v2/musesteamer/images/generations", + ) + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, object], + model_params, + provider, + raise_exception=False, + ): + for key in ["api_base", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key=key)) + return False + + try: + model = provider.get_model(model_type, model_name, model_credential, **model_params) + model.check_auth() + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + if isinstance(e, AppApiException): + raise e + if raise_exception: + raise AppApiException( + ValidCode.valid_error.value, + gettext("Verification failed, please check whether the parameters are correct: {error}").format( + error=str(e) + ), + ) + return False + return True + + def encryption_dict(self, model: Dict[str, object]): + return {**model, "api_key": super().encryption(model.get("api_key", ""))} + + def get_model_params_setting_form(self, model_name): + return QianfanTTIModelParams() diff --git a/apps/models_provider/impl/qianfan_model_provider/credential/ttv.py b/apps/models_provider/impl/qianfan_model_provider/credential/ttv.py new file mode 100644 index 00000000000..3603a808b53 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/credential/ttv.py @@ -0,0 +1,70 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: ttv.py +@desc: 千帆视频生成模型凭据 +""" + +from typing import Dict, Any + +from django.utils.translation import gettext, gettext_lazy as _ + +from common import forms +from common.exception.app_exception import AppApiException +from common.forms import BaseForm, SliderField, TooltipLabel +from common.forms.switch_field import SwitchField +from models_provider.base_model_provider import BaseModelCredential, ValidCode + + +class QianfanVideoModelParams(BaseForm): + duration = SliderField( + TooltipLabel( + _("Video duration"), + _("The duration of the generated video in seconds. Only supported values are accepted."), + ), + required=False, + default_value=None, + _min=1, + _max=30, + _step=1, + precision=0, + ) + + watermark = SwitchField( + TooltipLabel(_("Watermark"), _("Whether the generated video contains a watermark")), + attrs={"active-value": True, "inactive-value": False}, + default_value=False, + ) + + prompt_extend = SwitchField( + TooltipLabel(_("Prompt extend"), _("Whether to use a large model to rewrite the prompt")), + attrs={"active-value": True, "inactive-value": False}, + default_value=True, + ) + + +class QianfanVideoModelCredential(BaseForm, BaseModelCredential): + api_base = forms.TextInputField("API URL", required=True, default_value="https://qianfan.baidubce.com/v2") + api_key = forms.PasswordInputField("API Key", required=True) + + def is_valid( + self, + model_type: str, + model_name, + model_credential: Dict[str, Any], + model_params, + provider, + raise_exception=False, + ): + for key in ["api_base", "api_key"]: + if key not in model_credential: + if raise_exception: + raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key=key)) + return False + return True + + def encryption_dict(self, model: Dict[str, object]): + return {**model, "api_key": super().encryption(model.get("api_key", ""))} + + def get_model_params_setting_form(self, model_name): + return QianfanVideoModelParams() diff --git a/apps/models_provider/impl/qianfan_model_provider/model/__init__.py b/apps/models_provider/impl/qianfan_model_provider/model/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/apps/models_provider/impl/qianfan_model_provider/model/embedding.py b/apps/models_provider/impl/qianfan_model_provider/model/embedding.py new file mode 100644 index 00000000000..3d1bd049988 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/embedding.py @@ -0,0 +1,50 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author:虎 +@file: embedding.py +@date:2024/10/17 16:48 +@desc: +""" + +from typing import Dict, List + +import openai + +from models_provider.base_model_provider import MaxKBBaseEmbeddingModel + + +class QianfanEmbeddings(MaxKBBaseEmbeddingModel): + """千帆 OpenAI 兼容向量接口(v2)""" + + model_name: str + + def supports_image_embedding(self) -> bool: + return False + + @staticmethod + def is_cache_model(): + return False + + def __init__(self, api_key: str, base_url: str, model_name: str): + self.client = openai.OpenAI(api_key=api_key, base_url=base_url).embeddings + self.model_name = model_name + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + return QianfanEmbeddings( + api_key=model_credential.get("api_key"), + model_name=model_name, + base_url=model_credential.get("api_base"), + ) + + def embed_query(self, text: str): + res = self.embed_documents([text]) + return res[0] + + def embed_documents( + self, + texts: List[str], + ) -> List[List[float]]: + res = self.client.create(input=texts, model=self.model_name, encoding_format="float") + return [e.embedding for e in res.data] diff --git a/apps/models_provider/impl/qianfan_model_provider/model/image.py b/apps/models_provider/impl/qianfan_model_provider/model/image.py new file mode 100644 index 00000000000..bd99d8a5534 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/image.py @@ -0,0 +1,27 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: image.py +@desc: 千帆视觉理解模型(v2 OpenAI 兼容接口 /v2/chat/completions) +""" + +from typing import Dict + +from models_provider.base_model_provider import MaxKBBaseModel +from models_provider.impl.base_chat_open_ai import BaseChatOpenAI + + +class QianfanVisionModel(MaxKBBaseModel, BaseChatOpenAI): + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + optional_params = MaxKBBaseModel.filter_optional_params(model_kwargs) + return QianfanVisionModel( + model=model_name, + openai_api_base=model_credential.get("api_base"), + openai_api_key=model_credential.get("api_key"), + extra_body=optional_params, + ) diff --git a/apps/models_provider/impl/qianfan_model_provider/model/llm.py b/apps/models_provider/impl/qianfan_model_provider/model/llm.py new file mode 100644 index 00000000000..e08c404ba18 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/llm.py @@ -0,0 +1,31 @@ +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: llm.py +@date:2023/11/10 17:45 +@desc: +""" + +from typing import Dict + +from models_provider.base_model_provider import MaxKBBaseModel +from models_provider.impl.base_chat_open_ai import BaseChatOpenAI + + +class QianfanChatModel(MaxKBBaseModel, BaseChatOpenAI): + """千帆 OpenAI 兼容接口(v2)""" + + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + optional_params = MaxKBBaseModel.filter_optional_params(model_kwargs) + return QianfanChatModel( + model=model_name, + openai_api_base=model_credential.get("api_base"), + openai_api_key=model_credential.get("api_key"), + extra_body=optional_params, + ) diff --git a/apps/models_provider/impl/qianfan_model_provider/model/reranker.py b/apps/models_provider/impl/qianfan_model_provider/model/reranker.py new file mode 100644 index 00000000000..430d1f0f247 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/reranker.py @@ -0,0 +1,60 @@ +from typing import Sequence, Optional, Dict + +import requests +from langchain_core.callbacks import Callbacks +from langchain_core.documents import BaseDocumentCompressor, Document + +from models_provider.base_model_provider import MaxKBBaseModel + + +class QfBgeReranker(MaxKBBaseModel, BaseDocumentCompressor): + api_key: str + api_url: str + model: str + params: dict + top_n: int = 3 + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.api_key = kwargs.get("api_key") + self.model = kwargs.get("model") + self.params = kwargs.get("params", {}) + self.api_url = kwargs.get("api_url") + self.top_n = self.params.get("top_n", 3) + + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + return QfBgeReranker( + model=model_name, + api_key=model_credential.get("api_key"), + api_url=model_credential.get("api_url"), + params=model_kwargs, + ) + + def compress_documents( + self, documents: Sequence[Document], query: str, callbacks: Optional[Callbacks] = None + ) -> Sequence[Document]: + if not documents: + return [] + + texts = [doc.page_content for doc in documents] + + headers = {"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"} + top_n = min(self.top_n, len(texts)) + payload = {"model": self.model, "query": query, "documents": texts, "top_n": top_n} + + response = requests.post(f"{self.api_url}/rerank", json=payload, headers=headers) + + if response.status_code != 200: + raise RuntimeError(f"千帆 API 请求失败:{response.text}") + + res = response.json() + + return [ + Document(page_content=item.get("document", ""), metadata={"relevance_score": item.get("relevance_score")}) + for item in res.get("results", []) + ] diff --git a/apps/models_provider/impl/qianfan_model_provider/model/tti.py b/apps/models_provider/impl/qianfan_model_provider/model/tti.py new file mode 100644 index 00000000000..d4b27890970 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/tti.py @@ -0,0 +1,74 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: tti.py +@desc: 千帆文生图通用模型。api_base 即完整请求地址,直接使用。 +""" + +from typing import Dict + +import requests + +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import MaxKBBaseModel +from models_provider.impl.base_tti import BaseTextToImage + + +class QianfanTextToImage(MaxKBBaseModel, BaseTextToImage): + api_key: str + api_base: str + model_name: str + params: dict = {} + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.api_key = kwargs.get("api_key") + self.api_base = kwargs.get("api_base") + self.model_name = kwargs.get("model_name") + self.params = kwargs.get("params", {}) or {} + self._session = requests.Session() + self._session.headers.update({"Authorization": f"Bearer {self.api_key}"}) + + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + optional_params = {"params": {}} + for key, value in model_kwargs.items(): + if key not in ["model_id", "use_local", "streaming"]: + optional_params["params"][key] = value + return QianfanTextToImage( + model_name=model_name, + api_key=model_credential.get("api_key"), + api_base=model_credential.get("api_base"), + **optional_params, + ) + + def check_auth(self): + self.generate_image("a green grass field with a blue sky") + + def generate_image(self, prompt: str, negative_prompt: str = None): + payload = {"model": self.model_name, "prompt": prompt} + if negative_prompt: + payload["negative_prompt"] = negative_prompt + payload.update(self.params) + + try: + response = self._session.post(self.api_base, json=payload) + response.raise_for_status() + file_urls = [] + for item in response.json().get("data", []): + if not isinstance(item, dict): + continue + url = item.get("url") or item.get("b64_json") + if not url: + continue + if "://" not in url: + url = f"data:image/png;base64,{url}" + file_urls.append(url) + return file_urls + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + raise e diff --git a/apps/models_provider/impl/qianfan_model_provider/model/ttv.py b/apps/models_provider/impl/qianfan_model_provider/model/ttv.py new file mode 100644 index 00000000000..5137cdc80b0 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/model/ttv.py @@ -0,0 +1,124 @@ +# coding=utf-8 +""" +@project: MaxKB +@file: ttv.py +@desc: 千帆视频生成模型(蒸汽机 Air,异步任务式接口 /video/generations) +""" + +import time +from typing import ClassVar, Dict + +import requests + +from common.utils.logger import maxkb_logger +from models_provider.base_model_provider import MaxKBBaseModel +from models_provider.base_ttv import BaseGenerationVideo + + +class QianfanVideoModel(MaxKBBaseModel, BaseGenerationVideo): + api_key: str + api_base: str + model_name: str + params: dict = {} + + REQUEST_TIMEOUT: ClassVar[tuple] = (10, 120) + MAX_POLL_ATTEMPTS: ClassVar[int] = 180 + POLL_INTERVAL: ClassVar[int] = 5 + SUCCESS_STATUSES: ClassVar[frozenset] = frozenset({"succeeded"}) + FAIL_STATUSES: ClassVar[frozenset] = frozenset({"failed"}) + # 任务失败时用于提取错误信息的字段 + ERROR_KEYS: ClassVar[tuple] = ("error_msg", "error", "message", "msg", "detail", "description") + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.api_key = kwargs.get("api_key") + self.api_base = kwargs.get("api_base") + self.model_name = kwargs.get("model_name") + self.params = kwargs.get("params", {}) or {} + # 需要关闭默认的 Authorization 头污染 + self._session = requests.Session() + self._session.headers.update({"Authorization": f"Bearer {self.api_key}"}) + + @staticmethod + def is_cache_model(): + return False + + @staticmethod + def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs): + optional_params = {"params": {}} + for key, value in model_kwargs.items(): + if key not in ["model_id", "use_local", "streaming"]: + optional_params["params"][key] = value + return QianfanVideoModel( + model_name=model_name, + api_key=model_credential.get("api_key"), + api_base=model_credential.get("api_base", "https://qianfan.baidubce.com/v2"), + **optional_params, + ) + + def check_auth(self): + return True + + def _base_url(self): + """接口路径为 /video/generations(无 v2 前缀),去掉 api_base 末尾的 /v2。""" + base = self.api_base.rstrip("/") + if base.endswith("/v2"): + base = base[:-3] + return base.rstrip("/") + + def _request(self, method: str, url: str, **kwargs) -> dict: + kwargs.setdefault("timeout", self.REQUEST_TIMEOUT) + response = self._session.request(method, url, **kwargs) + try: + response.raise_for_status() + except requests.exceptions.HTTPError as e: + detail = e.response.text if e.response is not None else str(e) + maxkb_logger.error(f"千帆视频接口请求失败: {detail}", exc_info=True) + raise RuntimeError(f"HTTP 请求失败: {detail}") from e + return response.json() + + @staticmethod + def _extract_error(data: dict) -> str: + for key in QianfanVideoModel.ERROR_KEYS: + value = data.get(key) + if value: + return str(value) + return str(data) + + def _wait_for_result(self, task_id: str) -> dict: + query_url = f"{self._base_url()}/video/generations" + for attempt in range(1, self.MAX_POLL_ATTEMPTS + 1): + response_data = self._request("GET", query_url, params={"task_id": task_id}) + status = response_data.get("status") + maxkb_logger.info(f"千帆视频任务状态 (尝试 {attempt}/{self.MAX_POLL_ATTEMPTS}): {status}") + if status in self.SUCCESS_STATUSES: + return response_data + if status in self.FAIL_STATUSES: + raise RuntimeError(f"视频生成失败: {self._extract_error(response_data)}") + time.sleep(self.POLL_INTERVAL) + raise RuntimeError(f"任务超时:经过 {self.MAX_POLL_ATTEMPTS} 次轮询后仍未完成") + + def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, last_frame_url=None, **kwargs): + content = [{"type": "text", "text": prompt}] + if first_frame_url: + content.append({"type": "image_url", "image_url": {"url": first_frame_url}}) + if not any(item.get("type") == "image_url" for item in content): + # 图生视频(musesteamer-air-i2v)必须包含图片信息 + maxkb_logger.warning("千帆视频生成:未提供图片,文生视频接口可能不支持该模型") + + payload = {"model": self.model_name, "content": content} + payload.update(self.params) + + maxkb_logger.info(f"提交千帆视频生成任务,模型: {self.model_name}") + response_data = self._request("POST", f"{self._base_url()}/video/generations", json=payload) + + task_id = response_data.get("task_id") + if not task_id: + raise RuntimeError(f"提交任务失败,未获取到 task_id: {response_data}") + + response_data = self._wait_for_result(task_id) + + video_url = (response_data.get("content") or {}).get("video_url") + if not video_url: + raise RuntimeError(f"任务成功但未获取到 video_url: {response_data}") + return video_url diff --git a/apps/models_provider/impl/qianfan_model_provider/qianfan_model_provider.py b/apps/models_provider/impl/qianfan_model_provider/qianfan_model_provider.py new file mode 100644 index 00000000000..39f675fed56 --- /dev/null +++ b/apps/models_provider/impl/qianfan_model_provider/qianfan_model_provider.py @@ -0,0 +1,101 @@ +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: qianfan_model_provider.py +@date:2023/10/31 16:19 +@desc: +""" + +import os + +from common.utils.common import get_file_content +from models_provider.base_model_provider import ( + ModelProvideInfo, + ModelTypeConst, + ModelInfo, + IModelProvider, + ModelInfoManage, +) +from models_provider.impl.qianfan_model_provider.credential.embedding import QianfanEmbeddingCredential +from models_provider.impl.qianfan_model_provider.credential.image import QianfanImageModelCredential +from models_provider.impl.qianfan_model_provider.credential.llm import QianfanLLMModelCredential +from models_provider.impl.qianfan_model_provider.credential.reranker import QfRerankerCredential +from models_provider.impl.qianfan_model_provider.credential.tti import QianfanTextToImageModelCredential +from models_provider.impl.qianfan_model_provider.credential.ttv import QianfanVideoModelCredential +from models_provider.impl.qianfan_model_provider.model.embedding import QianfanEmbeddings +from models_provider.impl.qianfan_model_provider.model.image import QianfanVisionModel +from models_provider.impl.qianfan_model_provider.model.llm import QianfanChatModel +from models_provider.impl.qianfan_model_provider.model.tti import QianfanTextToImage +from models_provider.impl.qianfan_model_provider.model.ttv import QianfanVideoModel +from maxkb.conf import PROJECT_DIR +from django.utils.translation import gettext as _ + +from models_provider.impl.qianfan_model_provider.model.reranker import QfBgeReranker + +qianfan_llm_model_credential = QianfanLLMModelCredential() +qianfan_image_model_credential = QianfanImageModelCredential() +qianfan_tti_model_credential = QianfanTextToImageModelCredential() +qianfan_video_model_credential = QianfanVideoModelCredential() +qianfan_embedding_credential = QianfanEmbeddingCredential() +qf_reranker_credential = QfRerankerCredential() +model_info_list = [ + ModelInfo("ernie-5.1", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel), + ModelInfo("ernie-5.0", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel), + ModelInfo("ernie-4.5-turbo-128k", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel), + ModelInfo("deepseek-v4-pro", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel), + ModelInfo("ernie-4.5-turbo-32k", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel), +] +image_model_info_list = [ + ModelInfo("qwen2.5-vl-7b-instruct", "", ModelTypeConst.IMAGE, qianfan_image_model_credential, QianfanVisionModel), + ModelInfo("ernie-4.5-vl-28b-a3b", "", ModelTypeConst.IMAGE, qianfan_image_model_credential, QianfanVisionModel), +] +tti_model_info_list = [ + ModelInfo("musesteamer-air-image", "", ModelTypeConst.TTI, qianfan_tti_model_credential, QianfanTextToImage), + ModelInfo("qwen-image", "", ModelTypeConst.TTI, qianfan_tti_model_credential, QianfanTextToImage), + ModelInfo("ernie-image-turbo", "", ModelTypeConst.TTI, qianfan_tti_model_credential, QianfanTextToImage), +] +itv_model_info_list = [ + ModelInfo("musesteamer-air-i2v", "", ModelTypeConst.ITV, qianfan_video_model_credential, QianfanVideoModel), +] +embedding_model_info_list = [ + ModelInfo("Embedding-V1", "", ModelTypeConst.EMBEDDING, qianfan_embedding_credential, QianfanEmbeddings), + ModelInfo("bge-large-zh", "", ModelTypeConst.EMBEDDING, qianfan_embedding_credential, QianfanEmbeddings), +] +rerank_model_info_list = [ + ModelInfo("bce-reranker-base", "", ModelTypeConst.RERANKER, qf_reranker_credential, QfBgeReranker), +] +model_info_manage = ( + ModelInfoManage.builder() + .append_model_info_list(model_info_list) + .append_default_model_info( + ModelInfo("ernie-5.1", "", ModelTypeConst.LLM, qianfan_llm_model_credential, QianfanChatModel) + ) + .append_model_info_list(image_model_info_list) + .append_default_model_info(image_model_info_list[0]) + .append_model_info_list(tti_model_info_list) + .append_default_model_info(tti_model_info_list[0]) + .append_model_info_list(itv_model_info_list) + .append_default_model_info(itv_model_info_list[0]) + .append_model_info_list(embedding_model_info_list) + .append_default_model_info(embedding_model_info_list[0]) + .append_model_info_list(rerank_model_info_list) + .append_default_model_info(rerank_model_info_list[0]) + .build() +) + + +class QianfanModelProvider(IModelProvider): + def get_model_info_manage(self): + return model_info_manage + + def get_model_provide_info(self): + return ModelProvideInfo( + provider="model_qianfan_provider", + name=_("Thousand sails large model"), + icon=get_file_content( + os.path.join( + PROJECT_DIR, "apps", "models_provider", "impl", "qianfan_model_provider", "icon", "azure_icon_svg" + ) + ), + ) diff --git a/apps/models_provider/migrations/0002_rename_wenxin_provider_to_qianfan.py b/apps/models_provider/migrations/0002_rename_wenxin_provider_to_qianfan.py new file mode 100644 index 00000000000..33b974a5845 --- /dev/null +++ b/apps/models_provider/migrations/0002_rename_wenxin_provider_to_qianfan.py @@ -0,0 +1,24 @@ +from django.db import migrations + +OLD_PROVIDER = "model_wenxin_provider" +NEW_PROVIDER = "model_qianfan_provider" + + +def forwards(apps, schema_editor): + Model = apps.get_model("models_provider", "Model") + Model.objects.filter(provider=OLD_PROVIDER).update(provider=NEW_PROVIDER) + + +def backwards(apps, schema_editor): + Model = apps.get_model("models_provider", "Model") + Model.objects.filter(provider=NEW_PROVIDER).update(provider=OLD_PROVIDER) + + +class Migration(migrations.Migration): + dependencies = [ + ("models_provider", "0001_initial"), + ] + + operations = [ + migrations.RunPython(forwards, backwards), + ]