Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
# coding=utf-8
"""
@project: MaxKB
@desc: Tencent Tokenhub ASR sync_transcribe credential (model: wand-asr-v1 / hy-asr-3.0-preview)
"""

from django.utils.translation import gettext_lazy as _, gettext

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 TencentTokenhubSTTModelParams(BaseForm):
source = forms.SingleSelect(
label=TooltipLabel(_("Recognition language"), _("Recognition language: zh / en, auto detected when omitted")),
text_field="value",
value_field="value",
option_list=[
{"value": "", "label": _("Auto detect")},
{"value": "zh", "label": _("Chinese")},
{"value": "en", "label": _("English")},
],
required=False,
default_value="",
)
voice_encode_format = forms.SingleSelect(
label=TooltipLabel(_("Audio encoding"), _("pcm / wav / ogg / mp3, auto detected when omitted")),
text_field="value",
value_field="value",
option_list=[
{"value": "", "label": _("Auto")},
{"value": "pcm", "label": "pcm"},
{"value": "wav", "label": "wav"},
{"value": "ogg", "label": "ogg"},
{"value": "mp3", "label": "mp3"},
],
required=False,
default_value="",
)


class TencentTokenhubSTTModelCredential(BaseForm, BaseModelCredential):
def is_valid(self, model_type, model_name, model_credential, model_params, provider, raise_exception=False):
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,
gettext("{model_type} Model type is not supported").format(model_type=model_type),
)
if "api_key" not in model_credential:
if raise_exception:
raise AppApiException(ValidCode.valid_error.value, gettext("{key} is required").format(key="api_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 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):
return {**model, "api_key": super().encryption(model.get("api_key", ""))}

base_url = forms.TextInputField(
label=TooltipLabel(_("API URL"), _("Tokenhub sync_transcribe endpoint")),
required=False,
default_value="https://tokenhub.tencentmaas.com/v1/wand/asrproxy/sync_transcribe",
)
api_key = forms.PasswordInputField(_("API Key"), required=True)

def get_model_params_setting_form(self, model_name):
return TencentTokenhubSTTModelParams()
73 changes: 72 additions & 1 deletion apps/models_provider/impl/tencent_model_provider/model/stt.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
import json
import os
import traceback
from typing import Dict

import requests
from typing import Dict, Optional

from tencentcloud.asr.v20190614 import asr_client, models
from tencentcloud.common import credential
Expand Down Expand Up @@ -82,3 +84,72 @@ def speech_to_text(self, audio_file):
except TencentCloudSDKException as err:
maxkb_logger.error(f":Error: {str(err)}: {traceback.format_exc()}")
raise err


DEFAULT_WAND_BASE_URL = 'https://tokenhub.tencentmaas.com/v1/wand/asrproxy/sync_transcribe'


class TencentWandSpeechToText(MaxKBBaseModel, BaseSpeechToText):
api_key: str
model: str
params: dict
base_url: Optional[str] = DEFAULT_WAND_BASE_URL

def __init__(self, **kwargs):
super().__init__(**kwargs)
self.api_key = kwargs.get('api_key')
self.model = kwargs.get('model')
self.params = kwargs.get('params') or {}
self.base_url = kwargs.get('base_url') or DEFAULT_WAND_BASE_URL

@staticmethod
def is_cache_model():
return False

@staticmethod
def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs):
instance_kwargs = {
"api_key": model_credential.get('api_key'),
"model": model_name,
"params": model_kwargs,
**model_kwargs,
}
base_url = model_credential.get('base_url')
if base_url:
instance_kwargs["base_url"] = base_url
return TencentWandSpeechToText(**instance_kwargs)

def check_auth(self):
cwd = os.path.dirname(os.path.abspath(__file__))
with open(f'{cwd}/iat_mp3_16k.mp3', 'rb') as f:
self.speech_to_text(f)

def speech_to_text(self, audio_file):
try:
payload = {"model": self.model}
# 仅使用上传音频文件的 base64 data,不提供 input_url 兜底
audio_data = audio_file.read()
payload["data"] = base64.b64encode(audio_data).decode('utf-8')
for key in ("source", "voice_encode_format"):
if self.params.get(key):
payload[key] = self.params[key]

headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
response = requests.post(self.base_url, headers=headers, json=payload, timeout=300)
response.raise_for_status()
result = response.json()
if result.get("status") != "completed":
maxkb_logger.error(f"WAND ASR task not completed: {result}")
raise Exception(f"WAND ASR task not completed: {result}")
output = result.get("output") or {}
text = output.get("text")
if not text:
sentences = output.get("sentences") or []
text = " ".join([s.get('text', '') for s in sentences if s.get('text')])
return text
except Exception as e:
maxkb_logger.error(f"WAND ASR Error: {str(e)}: {traceback.format_exc()}")
raise e
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,12 @@
from models_provider.impl.tencent_model_provider.credential.image import TencentVisionModelCredential
from models_provider.impl.tencent_model_provider.credential.llm import TencentLLMModelCredential
from models_provider.impl.tencent_model_provider.credential.stt import TencentSTTModelCredential
from models_provider.impl.tencent_model_provider.credential.tokenhub_stt import TencentTokenhubSTTModelCredential
from models_provider.impl.tencent_model_provider.credential.tti import TencentTTIModelCredential
from models_provider.impl.tencent_model_provider.model.embedding import TencentEmbeddingModel
from models_provider.impl.tencent_model_provider.model.image import TencentVision
from models_provider.impl.tencent_model_provider.model.llm import TencentModel
from models_provider.impl.tencent_model_provider.model.stt import TencentSpeechToText
from models_provider.impl.tencent_model_provider.model.stt import TencentSpeechToText, TencentWandSpeechToText
from models_provider.impl.tencent_model_provider.model.tti import TencentTextToImageModel
from maxkb.conf import PROJECT_DIR
from django.utils.translation import gettext as _
Expand Down Expand Up @@ -78,6 +79,18 @@ def _initialize_model_info():
ModelTypeConst.STT,
TencentSTTModelCredential,
TencentSpeechToText),
_create_model_info(
'wand-asr-v1',
_('Tencent WAND ASR online recognition, supports long audio/video. Outputs text with sentence-level timestamps.'),
ModelTypeConst.STT,
TencentTokenhubSTTModelCredential,
TencentWandSpeechToText),
_create_model_info(
'hy-asr-3.0-preview',
_('Tencent Hunyuan ASR 3.0 (preview) online recognition. Accepts the same request/response format as wand-asr-v1.'),
ModelTypeConst.STT,
TencentTokenhubSTTModelCredential,
TencentWandSpeechToText),
]

tencent_embedding_model_info = _create_model_info(
Expand Down Expand Up @@ -125,6 +138,12 @@ def __init__(self):
def get_model_info_manage(self):
return self._model_info_manage

def get_model(self, model_type, model_name, model_credential, **model_kwargs):
# STT 模型:模型名不以 asr- 开头的一律走 Tencent Tokenhub WAND 识别
if model_type == ModelTypeConst.STT.name and not model_name.startswith('asr-'):
return TencentWandSpeechToText.new_instance(model_type, model_name, model_credential, **model_kwargs)
return super().get_model(model_type, model_name, model_credential, **model_kwargs)

def get_model_provide_info(self):
icon_path = _get_tencent_icon_path()
icon_data = get_file_content(icon_path)
Expand Down
Loading