mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 07:03:53 +08:00
Merge pull request #3157 from xinnan-tech/fix-streaming-tts-original-text
Fix streaming tts original text
This commit is contained in:
@@ -124,6 +124,8 @@ class TTSProvider(TTSProviderBase):
|
||||
)
|
||||
|
||||
if message.sentence_type == SentenceType.FIRST:
|
||||
# 重置流式处理状态
|
||||
self.reset_stream_state()
|
||||
# 初始化会话
|
||||
try:
|
||||
if not getattr(self.conn, "sentence_id", None):
|
||||
@@ -194,22 +196,24 @@ class TTSProvider(TTSProviderBase):
|
||||
|
||||
# 过滤Markdown
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
if self._correct_words_pattern:
|
||||
filtered_text = self._correct_words_pattern.sub(lambda m: self.correct_words[m.group(0)], filtered_text)
|
||||
|
||||
if filtered_text:
|
||||
# 发送continue-task消息
|
||||
continue_task_message = {
|
||||
"header": {
|
||||
"action": "continue-task",
|
||||
"task_id": self.conn.sentence_id,
|
||||
"streaming": "duplex",
|
||||
},
|
||||
"payload": {"input": {"text": filtered_text}},
|
||||
}
|
||||
# 使用滑动窗口匹配处理跨分片的替换词
|
||||
confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
|
||||
|
||||
await self.ws.send(json.dumps(continue_task_message))
|
||||
self.last_active_time = time.time()
|
||||
# 发送每个确定的文本片段
|
||||
for txt in confirmed_texts:
|
||||
if txt and self.ws:
|
||||
continue_task_message = {
|
||||
"header": {
|
||||
"action": "continue-task",
|
||||
"task_id": self.conn.sentence_id,
|
||||
"streaming": "duplex",
|
||||
},
|
||||
"payload": {"input": {"text": txt}},
|
||||
}
|
||||
await self.ws.send(json.dumps(continue_task_message))
|
||||
self.last_active_time = time.time()
|
||||
return
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||
|
||||
@@ -233,6 +233,8 @@ class TTSProvider(TTSProviderBase):
|
||||
)
|
||||
|
||||
if message.sentence_type == SentenceType.FIRST:
|
||||
# 重置流式处理状态
|
||||
self.reset_stream_state()
|
||||
# 初始化参数
|
||||
try:
|
||||
logger.bind(tag=TAG).debug("开始启动TTS会话...")
|
||||
@@ -295,21 +297,26 @@ class TTSProvider(TTSProviderBase):
|
||||
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
|
||||
return
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
if self._correct_words_pattern:
|
||||
filtered_text = self._correct_words_pattern.sub(lambda m: self.correct_words[m.group(0)], filtered_text)
|
||||
|
||||
if filtered_text:
|
||||
run_request = {
|
||||
"header": {
|
||||
"message_id": uuid.uuid4().hex,
|
||||
"task_id": self.task_id,
|
||||
"namespace": "FlowingSpeechSynthesizer",
|
||||
"name": "RunSynthesis",
|
||||
"appkey": self.appkey,
|
||||
},
|
||||
"payload": {"text": filtered_text},
|
||||
}
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
self.last_active_time = time.time()
|
||||
# 使用滑动窗口匹配处理跨分片的替换词
|
||||
confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
|
||||
|
||||
# 发送每个确定的文本片段
|
||||
for txt in confirmed_texts:
|
||||
if txt and self.ws:
|
||||
run_request = {
|
||||
"header": {
|
||||
"message_id": uuid.uuid4().hex,
|
||||
"task_id": self.task_id,
|
||||
"namespace": "FlowingSpeechSynthesizer",
|
||||
"name": "RunSynthesis",
|
||||
"appkey": self.appkey,
|
||||
},
|
||||
"payload": {"text": txt},
|
||||
}
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
self.last_active_time = time.time()
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -56,11 +56,28 @@ class TTSProviderBase(ABC):
|
||||
if self.correct_words:
|
||||
# 按key长度降序排列,长的先匹配,避免短词部分干扰
|
||||
sorted_keys = sorted(self.correct_words.keys(), key=len, reverse=True)
|
||||
pattern_str = '|'.join(re.escape(k) for k in sorted_keys)
|
||||
pattern_str = "|".join(re.escape(k) for k in sorted_keys)
|
||||
self._correct_words_pattern = re.compile(pattern_str)
|
||||
# 构建反向替换正则,用于将TTS服务返回的替换后文本还原为原始文本(字幕显示)
|
||||
reverse_map = {v: k for k, v in self.correct_words.items()}
|
||||
sorted_reverse_keys = sorted(reverse_map.keys(), key=len, reverse=True)
|
||||
reverse_pattern_str = "|".join(re.escape(k) for k in sorted_reverse_keys)
|
||||
self._reverse_words_pattern = re.compile(reverse_pattern_str)
|
||||
self._reverse_words_map = reverse_map
|
||||
# 流式滑动窗口:按首字分组的替换词字典,用于快速查找
|
||||
self._words_by_first_char = {}
|
||||
for key in sorted_keys: # 使用已按长度降序排列的keys,确保长词优先匹配
|
||||
first_char = key[0] if key else ""
|
||||
if first_char not in self._words_by_first_char:
|
||||
self._words_by_first_char[first_char] = []
|
||||
self._words_by_first_char[first_char].append(key)
|
||||
else:
|
||||
self._correct_words_pattern = None
|
||||
self._reverse_words_pattern = None
|
||||
self._reverse_words_map = None
|
||||
|
||||
# 流式滑动窗口:待匹配的缓存文本
|
||||
self._pending_prefix = ""
|
||||
self.tts_text_buff = []
|
||||
self.punctuations = (
|
||||
"。",
|
||||
@@ -339,6 +356,13 @@ class TTSProviderBase(ABC):
|
||||
if sentence_id in self._sentence_text_map:
|
||||
del self._sentence_text_map[sentence_id]
|
||||
|
||||
def _restore_original_text(self, text):
|
||||
if not self._reverse_words_pattern or not text:
|
||||
return text
|
||||
return self._reverse_words_pattern.sub(
|
||||
lambda m: self._reverse_words_map[m.group(0)], text
|
||||
)
|
||||
|
||||
# 这里默认是非流式的处理方式
|
||||
# 流式处理方式请在子类中重写
|
||||
def tts_text_priority_thread(self):
|
||||
@@ -549,3 +573,64 @@ class TTSProviderBase(ABC):
|
||||
if config_key in config:
|
||||
val = convert_percentage_to_range(config[config_key], min_val, max_val, base_val)
|
||||
setattr(self, attr_name, transform(val) if transform else val)
|
||||
|
||||
def _match_stream_text(self, text):
|
||||
"""流式文本滑动窗口匹配,用于处理跨分片的替换词
|
||||
|
||||
Args:
|
||||
text: 输入的文本片段
|
||||
|
||||
Returns:
|
||||
tuple: (确定的文本列表, 剩余待匹配的前缀)
|
||||
"""
|
||||
if not self.correct_words or not text:
|
||||
return [text] if text else [], ""
|
||||
|
||||
result = []
|
||||
pending = self._pending_prefix
|
||||
i = 0
|
||||
|
||||
while i < len(text):
|
||||
char = text[i]
|
||||
|
||||
# 尝试:pending + 当前字符 是否能匹配替换词
|
||||
test_text = pending + char
|
||||
|
||||
matched = False
|
||||
# 遍历可能匹配的替换词
|
||||
candidates = self._words_by_first_char.get(pending[0], []) if pending else self._words_by_first_char.get(char, [])
|
||||
for key in candidates:
|
||||
if test_text == key:
|
||||
# 完整匹配,替换后发送
|
||||
result.append(self.correct_words[key])
|
||||
pending = ""
|
||||
matched = True
|
||||
break
|
||||
elif key.startswith(test_text):
|
||||
# 是替换词的前缀,继续等待
|
||||
pending = test_text
|
||||
matched = True
|
||||
break
|
||||
|
||||
if matched:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
# 没有匹配到更长的词,pending 的内容确定可以发送
|
||||
if pending:
|
||||
result.append(pending)
|
||||
pending = ""
|
||||
|
||||
# 检查当前字符是否是某个替换词的开头
|
||||
if char in self._words_by_first_char:
|
||||
pending = char
|
||||
else:
|
||||
result.append(char)
|
||||
|
||||
i += 1
|
||||
|
||||
return result, pending
|
||||
|
||||
def reset_stream_state(self):
|
||||
"""重置流式处理状态,用于会话开始时清理残留状态"""
|
||||
self._pending_prefix = ""
|
||||
|
||||
@@ -300,6 +300,8 @@ class TTSProvider(TTSProviderBase):
|
||||
)
|
||||
|
||||
if message.sentence_type == SentenceType.FIRST:
|
||||
# 重置流式处理状态
|
||||
self.reset_stream_state()
|
||||
# 初始化参数
|
||||
try:
|
||||
if not getattr(self.conn, "sentence_id", None):
|
||||
@@ -370,12 +372,16 @@ class TTSProvider(TTSProviderBase):
|
||||
|
||||
# 过滤Markdown
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
if self._correct_words_pattern:
|
||||
filtered_text = self._correct_words_pattern.sub(lambda m: self.correct_words[m.group(0)], filtered_text)
|
||||
|
||||
if filtered_text:
|
||||
# 发送文本
|
||||
await self.send_text(self.voice, filtered_text, self.conn.sentence_id)
|
||||
# 使用滑动窗口匹配处理跨分片的替换词
|
||||
confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
|
||||
|
||||
# 发送每个确定的文本片段
|
||||
for txt in confirmed_texts:
|
||||
if txt and self.ws:
|
||||
await self.send_text(self.voice, txt, self.conn.sentence_id)
|
||||
|
||||
return
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||
@@ -514,9 +520,9 @@ class TTSProvider(TTSProviderBase):
|
||||
json_data = json.loads(res.payload.decode("utf-8"))
|
||||
self.tts_text = json_data.get("text", "")
|
||||
logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}")
|
||||
self.tts_audio_queue.put(
|
||||
(SentenceType.FIRST, [], self.tts_text)
|
||||
)
|
||||
# 将TTS服务返回的替换后文本还原为原始文本,用于字幕显示
|
||||
display_text = self._restore_original_text(self.tts_text)
|
||||
self.tts_audio_queue.put((SentenceType.FIRST, [], display_text))
|
||||
elif (
|
||||
res.optional.event == EVENT_TTSResponse
|
||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||
|
||||
@@ -168,6 +168,8 @@ class TTSProvider(TTSProviderBase):
|
||||
)
|
||||
|
||||
if message.sentence_type == SentenceType.FIRST:
|
||||
# 重置流式处理状态
|
||||
self.reset_stream_state()
|
||||
# 重置序列号
|
||||
self.text_seq = 0
|
||||
# 增加序列号
|
||||
@@ -245,12 +247,17 @@ class TTSProvider(TTSProviderBase):
|
||||
return
|
||||
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
if self._correct_words_pattern:
|
||||
filtered_text = self._correct_words_pattern.sub(lambda m: self.correct_words[m.group(0)], filtered_text)
|
||||
|
||||
if filtered_text:
|
||||
# 发送文本合成请求
|
||||
run_request = self._build_base_request(status=1,text=filtered_text)
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
# 使用滑动窗口匹配处理跨分片的替换词
|
||||
confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
|
||||
|
||||
# 发送每个确定的文本片段
|
||||
for txt in confirmed_texts:
|
||||
if txt and self.ws:
|
||||
# 发送文本合成请求
|
||||
run_request = self._build_base_request(status=1, text=txt)
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
|
||||
Reference in New Issue
Block a user