mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 09:03:54 +08:00
update: 增加豆包流式多语种识别
This commit is contained in:
@@ -346,6 +346,11 @@ ASR:
|
||||
# 热词、替换词使用流程:https://www.volcengine.com/docs/6561/155738
|
||||
boosting_table_name: (选填)你的热词文件名称
|
||||
correct_table_name: (选填)你的替换词文件名称
|
||||
# 是否开启多语种识别模式
|
||||
enable_multilingual: False
|
||||
# 多语种识别当该键为空时,该模型支持中英文、上海话、闽南语,四川、陕西、粤语识别。当将其设置为特定键时,它可以识别指定语言。
|
||||
# 详细语言列表参考 https://www.volcengine.com/docs/6561/1354869
|
||||
# language: zh-cn
|
||||
# 静音判定时长(ms),默认200ms
|
||||
end_window_size: 200
|
||||
output_dir: tmp/
|
||||
@@ -496,11 +501,11 @@ ASR:
|
||||
# 可选参数
|
||||
disfluency_removal_enabled: false # 是否过滤语气词(如"嗯"、"啊"等)
|
||||
semantic_punctuation_enabled: false # 语义断句(true:会议场景,准确;false:VAD断句,交互场景,低延迟)
|
||||
max_sentence_silence: 800 # VAD断句静音时长阈值(ms),范围200-6000,仅VAD断句时生效
|
||||
max_sentence_silence: 200 # VAD断句静音时长阈值(ms),范围200-6000,仅VAD断句时生效
|
||||
multi_threshold_mode_enabled: false # 防止VAD断句切割过长,仅VAD断句时生效
|
||||
punctuation_prediction_enabled: true # 是否自动添加标点符号
|
||||
heartbeat: false # 是否开启长连接心跳(持续静音下保持连接)
|
||||
inverse_text_normalization_enabled: true # 是否开启ITN(中文数字转阿拉伯数字)
|
||||
# 热词定制文档地址:https://help.aliyun.com/zh/model-studio/custom-hot-words?
|
||||
# vocabulary_id: vocab-xxx-24ee19fa8cfb4d52902170a0xxxxxxxx # 热词ID(可选)
|
||||
# language_hints: ["zh", "en"] # 指定语言(可选),支持zh、en、ja、yue、ko、de、fr、ru
|
||||
output_dir: tmp/
|
||||
|
||||
@@ -32,15 +32,15 @@ class ASRProvider(ASRProviderBase):
|
||||
self.format = config.get("format", "pcm")
|
||||
|
||||
# 可选参数
|
||||
self.vocabulary_id = config.get("vocabulary_id") # 热词ID
|
||||
self.disfluency_removal_enabled = config.get("disfluency_removal_enabled", False) # 过滤语气词
|
||||
self.language_hints = config.get("language_hints") # 语言提示,如 ["zh", "en"]
|
||||
self.semantic_punctuation_enabled = config.get("semantic_punctuation_enabled", False) # 语义断句
|
||||
self.max_sentence_silence = config.get("max_sentence_silence", 800) # VAD断句静音时长(ms)
|
||||
self.multi_threshold_mode_enabled = config.get("multi_threshold_mode_enabled", False) # 防止VAD断句切割过长
|
||||
self.punctuation_prediction_enabled = config.get("punctuation_prediction_enabled", True) # 标点符号预测
|
||||
self.heartbeat = config.get("heartbeat", False) # 长连接心跳
|
||||
self.inverse_text_normalization_enabled = config.get("inverse_text_normalization_enabled", True) # ITN
|
||||
self.vocabulary_id = config.get("vocabulary_id")
|
||||
self.disfluency_removal_enabled = config.get("disfluency_removal_enabled", False)
|
||||
self.language_hints = config.get("language_hints")
|
||||
self.semantic_punctuation_enabled = config.get("semantic_punctuation_enabled", False)
|
||||
max_sentence_silence = config.get("max_sentence_silence")
|
||||
self.max_sentence_silence = int(max_sentence_silence) if max_sentence_silence else 200
|
||||
self.multi_threshold_mode_enabled = config.get("multi_threshold_mode_enabled", False)
|
||||
self.punctuation_prediction_enabled = config.get("punctuation_prediction_enabled", True)
|
||||
self.inverse_text_normalization_enabled = config.get("inverse_text_normalization_enabled", True)
|
||||
|
||||
# WebSocket URL
|
||||
self.ws_url = "wss://dashscope.aliyuncs.com/api-ws/v1/inference"
|
||||
@@ -85,6 +85,10 @@ class ASRProvider(ASRProviderBase):
|
||||
async def _start_recognition(self, conn):
|
||||
"""开始识别会话"""
|
||||
try:
|
||||
# 如果为手动模式,设置超时时长为最大值
|
||||
if conn.client_listen_mode == "manual":
|
||||
self.max_sentence_silence = 6000
|
||||
|
||||
self.is_processing = True
|
||||
self.task_id = uuid.uuid4().hex
|
||||
|
||||
@@ -143,21 +147,21 @@ class ASRProvider(ASRProviderBase):
|
||||
"max_sentence_silence": self.max_sentence_silence,
|
||||
"multi_threshold_mode_enabled": self.multi_threshold_mode_enabled,
|
||||
"punctuation_prediction_enabled": self.punctuation_prediction_enabled,
|
||||
"heartbeat": self.heartbeat,
|
||||
"inverse_text_normalization_enabled": self.inverse_text_normalization_enabled,
|
||||
},
|
||||
"input": {}
|
||||
}
|
||||
}
|
||||
|
||||
# 添加可选参数
|
||||
if self.vocabulary_id:
|
||||
# 只有当模型名称以v2结尾时才添加vocabulary_id参数
|
||||
if self.model.lower().endswith("v2"):
|
||||
message["payload"]["parameters"]["vocabulary_id"] = self.vocabulary_id
|
||||
|
||||
if self.language_hints:
|
||||
message["payload"]["parameters"]["language_hints"] = self.language_hints
|
||||
|
||||
return message
|
||||
|
||||
async def _forward_results(self, conn):
|
||||
"""转发识别结果"""
|
||||
try:
|
||||
@@ -191,21 +195,10 @@ class ASRProvider(ASRProviderBase):
|
||||
output = payload.get("output", {})
|
||||
sentence = output.get("sentence", {})
|
||||
|
||||
if not sentence:
|
||||
continue
|
||||
|
||||
# 跳过心跳消息
|
||||
if sentence.get("heartbeat", False):
|
||||
continue
|
||||
|
||||
text = sentence.get("text", "")
|
||||
sentence_end = sentence.get("sentence_end", False)
|
||||
end_time = sentence.get("end_time")
|
||||
|
||||
# 只处理有文本的结果
|
||||
if not text:
|
||||
continue
|
||||
|
||||
# 判断是否为最终结果(sentence_end为True且end_time不为null)
|
||||
is_final = sentence_end and end_time is not None
|
||||
|
||||
@@ -272,8 +265,11 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
async def _send_stop_request(self):
|
||||
"""发送停止请求(用于手动模式停止录音)"""
|
||||
if self.asr_ws and self.task_id and self.is_processing:
|
||||
if self.asr_ws:
|
||||
try:
|
||||
# 先停止音频发送
|
||||
self.is_processing = False
|
||||
|
||||
logger.bind(tag=TAG).debug("收到停止请求,发送finish-task指令")
|
||||
await self._send_finish_task()
|
||||
except Exception as e:
|
||||
|
||||
@@ -33,7 +33,12 @@ class ASRProvider(ASRProviderBase):
|
||||
self.delete_audio_file = delete_audio_file
|
||||
|
||||
# 火山引擎ASR配置
|
||||
self.ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel"
|
||||
enable_multilingual = config.get("enable_multilingual", False)
|
||||
self.enable_multilingual = False if str(enable_multilingual).lower() == 'false' else True
|
||||
if self.enable_multilingual:
|
||||
self.ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel_nostream"
|
||||
else:
|
||||
self.ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel"
|
||||
self.uid = config.get("uid", "streaming_asr_service")
|
||||
self.workflow = config.get(
|
||||
"workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
|
||||
@@ -42,7 +47,8 @@ class ASRProvider(ASRProviderBase):
|
||||
self.format = config.get("format", "pcm")
|
||||
self.codec = config.get("codec", "pcm")
|
||||
self.rate = config.get("sample_rate", 16000)
|
||||
self.language = config.get("language", "zh-CN")
|
||||
# language参数仅在多语种模式(bigmodel_nostream)下有效
|
||||
self.language = config.get("language") if self.enable_multilingual else None
|
||||
self.bits = config.get("bits", 16)
|
||||
self.channel = config.get("channel", 1)
|
||||
self.auth_method = config.get("auth_method", "token")
|
||||
@@ -175,7 +181,8 @@ class ASRProvider(ASRProviderBase):
|
||||
utterances = payload["result"].get("utterances", [])
|
||||
# 检查duration和空文本的情况
|
||||
if (
|
||||
payload.get("audio_info", {}).get("duration", 0) > 2000
|
||||
not self.enable_multilingual # 注意:多语种模式不返回中间结果,需要等待最终结果
|
||||
and payload.get("audio_info", {}).get("duration", 0) > 2000
|
||||
and not utterances
|
||||
and not payload["result"].get("text")
|
||||
and conn.client_listen_mode != "manual"
|
||||
@@ -189,6 +196,10 @@ class ASRProvider(ASRProviderBase):
|
||||
|
||||
# 专门处理没有文本的识别结果(手动模式下可能已经识别完成但是没松按键)
|
||||
elif not payload["result"].get("text") and not utterances:
|
||||
# 多语种模式会持续返回空文本,直到最后返回完整结果,所以需要排除
|
||||
if self.enable_multilingual:
|
||||
continue
|
||||
|
||||
if conn.client_listen_mode == "manual" and conn.client_voice_stop and len(audio_data) > 0:
|
||||
logger.bind(tag=TAG).debug("消息结束收到停止信号,触发处理")
|
||||
await self.handle_voice_stop(conn, audio_data)
|
||||
@@ -299,12 +310,16 @@ class ASRProvider(ASRProviderBase):
|
||||
"format": self.format,
|
||||
"codec": self.codec,
|
||||
"rate": self.rate,
|
||||
"language": self.language,
|
||||
"bits": self.bits,
|
||||
"channel": self.channel,
|
||||
"sample_rate": self.rate,
|
||||
},
|
||||
}
|
||||
|
||||
# language参数仅在多语种模式下添加
|
||||
if self.enable_multilingual and self.language:
|
||||
req["audio"]["language"] = self.language
|
||||
|
||||
logger.bind(tag=TAG).debug(
|
||||
f"构造请求参数: {json.dumps(req, ensure_ascii=False)}"
|
||||
)
|
||||
|
||||
@@ -164,7 +164,7 @@ class TTSProvider(TTSProviderBase):
|
||||
self.authorization = config.get("authorization")
|
||||
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||
enable_ws_reuse_value = config.get("enable_ws_reuse", True)
|
||||
self.enable_ws_reuse = False if str(enable_ws_reuse_value).lower() in ('false', 'False') else True
|
||||
self.enable_ws_reuse = False if str(enable_ws_reuse_value).lower() == 'false' else True
|
||||
self.tts_text = ""
|
||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||
sample_rate=16000, channels=1, frame_size_ms=60
|
||||
|
||||
Reference in New Issue
Block a user