Merge pull request #3157 from xinnan-tech/fix-streaming-tts-original-text

Fix streaming tts original text
This commit is contained in:
wengzh
2026-05-07 15:36:56 +08:00
committed by GitHub
5 changed files with 149 additions and 40 deletions
@@ -124,6 +124,8 @@ class TTSProvider(TTSProviderBase):
) )
if message.sentence_type == SentenceType.FIRST: if message.sentence_type == SentenceType.FIRST:
# 重置流式处理状态
self.reset_stream_state()
# 初始化会话 # 初始化会话
try: try:
if not getattr(self.conn, "sentence_id", None): if not getattr(self.conn, "sentence_id", None):
@@ -194,22 +196,24 @@ class TTSProvider(TTSProviderBase):
# 过滤Markdown # 过滤Markdown
filtered_text = MarkdownCleaner.clean_markdown(text) 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: if filtered_text:
# 发送continue-task消息 # 使用滑动窗口匹配处理跨分片的替换词
continue_task_message = { confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
"header": {
"action": "continue-task",
"task_id": self.conn.sentence_id,
"streaming": "duplex",
},
"payload": {"input": {"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 return
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}") logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
@@ -233,6 +233,8 @@ class TTSProvider(TTSProviderBase):
) )
if message.sentence_type == SentenceType.FIRST: if message.sentence_type == SentenceType.FIRST:
# 重置流式处理状态
self.reset_stream_state()
# 初始化参数 # 初始化参数
try: try:
logger.bind(tag=TAG).debug("开始启动TTS会话...") logger.bind(tag=TAG).debug("开始启动TTS会话...")
@@ -295,21 +297,26 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本") logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
return return
filtered_text = MarkdownCleaner.clean_markdown(text) 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: if filtered_text:
run_request = { # 使用滑动窗口匹配处理跨分片的替换词
"header": { confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
"message_id": uuid.uuid4().hex,
"task_id": self.task_id, # 发送每个确定的文本片段
"namespace": "FlowingSpeechSynthesizer", for txt in confirmed_texts:
"name": "RunSynthesis", if txt and self.ws:
"appkey": self.appkey, run_request = {
}, "header": {
"payload": {"text": filtered_text}, "message_id": uuid.uuid4().hex,
} "task_id": self.task_id,
await self.ws.send(json.dumps(run_request)) "namespace": "FlowingSpeechSynthesizer",
self.last_active_time = time.time() "name": "RunSynthesis",
"appkey": self.appkey,
},
"payload": {"text": txt},
}
await self.ws.send(json.dumps(run_request))
self.last_active_time = time.time()
return return
except Exception as e: except Exception as e:
+86 -1
View File
@@ -56,11 +56,28 @@ class TTSProviderBase(ABC):
if self.correct_words: if self.correct_words:
# 按key长度降序排列,长的先匹配,避免短词部分干扰 # 按key长度降序排列,长的先匹配,避免短词部分干扰
sorted_keys = sorted(self.correct_words.keys(), key=len, reverse=True) 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) 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: else:
self._correct_words_pattern = None self._correct_words_pattern = None
self._reverse_words_pattern = None
self._reverse_words_map = None
# 流式滑动窗口:待匹配的缓存文本
self._pending_prefix = ""
self.tts_text_buff = [] self.tts_text_buff = []
self.punctuations = ( self.punctuations = (
"", "",
@@ -339,6 +356,13 @@ class TTSProviderBase(ABC):
if sentence_id in self._sentence_text_map: if sentence_id in self._sentence_text_map:
del self._sentence_text_map[sentence_id] 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): def tts_text_priority_thread(self):
@@ -549,3 +573,64 @@ class TTSProviderBase(ABC):
if config_key in config: if config_key in config:
val = convert_percentage_to_range(config[config_key], min_val, max_val, base_val) val = convert_percentage_to_range(config[config_key], min_val, max_val, base_val)
setattr(self, attr_name, transform(val) if transform else 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: if message.sentence_type == SentenceType.FIRST:
# 重置流式处理状态
self.reset_stream_state()
# 初始化参数 # 初始化参数
try: try:
if not getattr(self.conn, "sentence_id", None): if not getattr(self.conn, "sentence_id", None):
@@ -370,12 +372,16 @@ class TTSProvider(TTSProviderBase):
# 过滤Markdown # 过滤Markdown
filtered_text = MarkdownCleaner.clean_markdown(text) 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: 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 return
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(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")) json_data = json.loads(res.payload.decode("utf-8"))
self.tts_text = json_data.get("text", "") self.tts_text = json_data.get("text", "")
logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}") logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}")
self.tts_audio_queue.put( # 将TTS服务返回的替换后文本还原为原始文本,用于字幕显示
(SentenceType.FIRST, [], self.tts_text) display_text = self._restore_original_text(self.tts_text)
) self.tts_audio_queue.put((SentenceType.FIRST, [], display_text))
elif ( elif (
res.optional.event == EVENT_TTSResponse res.optional.event == EVENT_TTSResponse
and res.header.message_type == AUDIO_ONLY_RESPONSE and res.header.message_type == AUDIO_ONLY_RESPONSE
@@ -168,6 +168,8 @@ class TTSProvider(TTSProviderBase):
) )
if message.sentence_type == SentenceType.FIRST: if message.sentence_type == SentenceType.FIRST:
# 重置流式处理状态
self.reset_stream_state()
# 重置序列号 # 重置序列号
self.text_seq = 0 self.text_seq = 0
# 增加序列号 # 增加序列号
@@ -245,12 +247,17 @@ class TTSProvider(TTSProviderBase):
return return
filtered_text = MarkdownCleaner.clean_markdown(text) 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: if filtered_text:
# 发送文本合成请求 # 使用滑动窗口匹配处理跨分片的替换词
run_request = self._build_base_request(status=1,text=filtered_text) confirmed_texts, self._pending_prefix = self._match_stream_text(filtered_text)
await self.ws.send(json.dumps(run_request))
# 发送每个确定的文本片段
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 return
except Exception as e: except Exception as e: