Przeglądaj źródła

feat(tts): xiaozhi-style two-tier sentence segmentation + markdown cleaner + think tag filter

Adopting key patterns from xiaozhi-esp32-server:
- Two-tier TTS segmentation: first sentence uses commas as boundary,
  subsequent sentences use sentence-ending punctuation only
- _TTSSegmenter with processed_chars tracking (no destructive slicing)
- <think>/</think> tag filtering in LLM streaming output
- _clean_tts_text() markdown stripper for TTS input
- SYSTEM_PROMPT with identity constraint and anti-markdown rules
wenhongquan 3 tygodni temu
rodzic
commit
9788bcb353
1 zmienionych plików z 115 dodań i 69 usunięć
  1. 115 69
      asr_agent/conversation_worker.py

+ 115 - 69
asr_agent/conversation_worker.py

@@ -1,10 +1,10 @@
 """
-Full Pipeline Worker: VAD → ASR → LLM → TTS
+Full Pipeline Worker: VAD -> ASR -> LLM -> TTS
   VAD: dual-threshold energy-based + sliding window (xiaozhi-inspired)
   VP:  AEC/ANS/AGC handled by LiveKit WebRTC
   ASR: Qwen3-ASR (local model, batch on speech-end)
-  LLM: streaming via OpenAI-compatible API, incremental TTS
-  TTS: Mimo API (remote), per-sentence segmentation
+  LLM: streaming via OpenAI-compatible API (vLLM), <think> tag filtering
+  TTS: Mimo API (remote), xiaozhi-style two-tier sentence segmentation
 """
 
 from __future__ import annotations
@@ -34,7 +34,7 @@ from qwen_engine import QwenASREngine
 
 logger = logging.getLogger("worker")
 
-# ── Environment ──
+# ── Env ──
 LIVEKIT_URL = os.environ.get("LIVEKIT_URL", "ws://localhost:7880")
 ASR_MODEL = os.environ.get("ASR_MODEL_PATH", "Qwen/Qwen3-ASR-0.6B")
 VLLM_URL = os.environ.get("VLLM_URL", "http://127.0.0.1:8000/v1")
@@ -43,20 +43,18 @@ MIMO_KEY = os.environ.get("MIMO_KEY", "")
 
 # ── VAD ──
 VAD_THRESHOLD_HIGH = 0.015
-VAD_THRESHOLD_LOW = 0.005
-VAD_WINDOW_SIZE = 8
-VAD_VOICE_RATIO = 0.5
-MIN_SPEECH_S = 0.3
-MIN_SILENCE_S = 0.8
-MAX_SPEECH_S = 8.0
+VAD_THRESHOLD_LOW  = 0.005
+VAD_WINDOW_SIZE    = 8
+VAD_VOICE_RATIO    = 0.5
+MIN_SPEECH_S       = 0.3
+MIN_SILENCE_S      = 0.8
+MAX_SPEECH_S       = 8.0
 
 # ── LLM ──
-LLM_MAX_TOKENS = 128
+LLM_MAX_TOKENS  = 128
 LLM_TEMPERATURE = 0.7
-LLM_TIMEOUT = 20
+LLM_TIMEOUT     = 25
 
-# ── TTS ──
-TTS_CHUNK_PATTERN = re.compile(r"[。!?;\n]")
 SYSTEM_PROMPT = (
     "你是友好的中文语音助手,名字叫小智。"
     "回答简洁自然,2-3句即可,用口语化中文。"
@@ -64,8 +62,16 @@ SYSTEM_PROMPT = (
     "不要输出括号注释、不要使用英文缩写。"
 )
 
+# ── TTS segmentation (xiaozhi-style two-tier) ──
+FIRST_SENTENCE_PUNCT = ",、,。!?;::\n"
+SENTENCE_END_PUNCT   = "。!?!?\n"
+MIN_TTS_CHARS        = 2
+
+
+# ═══════════════════════════════════════════════════════════════
+#  JWT
+# ═══════════════════════════════════════════════════════════════
 
-# ── JWT ──
 def _token(room, ident):
     n = int(time.time())
     return jwt.encode({
@@ -76,9 +82,11 @@ def _token(room, ident):
     }, "secretsecretsecretsecretsecret12", algorithm="HS256")
 
 
-# ── LLM Streaming ──
+# ═══════════════════════════════════════════════════════════════
+#  LLM (streaming, <think> filtering)
+# ═══════════════════════════════════════════════════════════════
+
 async def _llm_stream(prompt: str, hist: list[dict]):
-    """Stream LLM tokens via SSE, yields (delta_text, is_final)."""
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     msgs.extend(hist)
     msgs.append({"role": "user", "content": prompt})
@@ -88,10 +96,10 @@ async def _llm_stream(prompt: str, hist: list[dict]):
             f"{VLLM_URL}/chat/completions",
             json={"model": LLM_MODEL, "messages": msgs,
                   "max_tokens": LLM_MAX_TOKENS, "temperature": LLM_TEMPERATURE,
-                  "stream": True},
+                  "stream": True,
+                  "chat_template_kwargs": {"enable_thinking": False}},
             timeout=aiohttp.ClientTimeout(total=LLM_TIMEOUT),
         ) as r:
-            full = ""
             async for line in r.content:
                 line = line.decode().strip()
                 if not line.startswith("data: "):
@@ -102,16 +110,23 @@ async def _llm_stream(prompt: str, hist: list[dict]):
                 try:
                     chunk = json.loads(data)
                     delta = chunk.get("choices", [{}])[0].get("delta", {})
-                    content = delta.get("content", "")
-                    if content:
-                        full += content
-                        yield content, False
+                    text = delta.get("content", "")
+                    if text:
+                        if "</think>" in text:
+                            text = text.split("</think>")[-1]
+                        if "<think>" in text:
+                            text = text.split("<think>")[0]
+                        if text.strip():
+                            yield text, False
                 except Exception:
                     continue
     yield "", True
 
 
-# ── TTS ──
+# ═══════════════════════════════════════════════════════════════
+#  TTS
+# ═══════════════════════════════════════════════════════════════
+
 async def _tts(text: str) -> np.ndarray | None:
     key = os.environ.get("MIMO_KEY", "")
     if not key:
@@ -133,7 +148,59 @@ async def _tts(text: str) -> np.ndarray | None:
                 return None
 
 
-# ── Enhanced VAD (dual-threshold + sliding window) ──
+def _clean_tts_text(text: str) -> str:
+    text = re.sub(r"\*{1,3}(.*?)\*{1,3}", r"\1", text)
+    text = re.sub(r"`{1,3}.*?`{1,3}", "", text)
+    text = re.sub(r"\[([^\]]+)]\([^)]+\)", r"\1", text)
+    text = re.sub(r"#{1,6}\s*", "", text)
+    text = re.sub(r"[>\-]\s", "", text)
+    text = re.sub(r"[\U0001F300-\U0001F9FF]", "", text)
+    return text.strip()
+
+
+# ═══════════════════════════════════════════════════════════════
+#  TTS Segmenter (xiaozhi-style: two-tier, processed_chars)
+# ═══════════════════════════════════════════════════════════════
+
+@dataclass
+class _TTSSegmenter:
+    buffer: str = ""
+    processed: int = 0
+    is_first: bool = True
+
+    def feed(self, text: str) -> list[str]:
+        self.buffer += text
+        segments = []
+        puncts = FIRST_SENTENCE_PUNCT if self.is_first else SENTENCE_END_PUNCT
+        unprocessed = self.buffer[self.processed:]
+
+        last_pos = -1
+        for ch in puncts:
+            pos = unprocessed.rfind(ch)
+            if pos > last_pos:
+                last_pos = pos
+
+        if last_pos >= 0:
+            seg_raw = unprocessed[:last_pos + 1]
+            seg_clean = _clean_tts_text(seg_raw)
+            if len(seg_clean) >= MIN_TTS_CHARS:
+                segments.append(seg_clean)
+                self.processed += len(seg_raw)
+                if self.is_first:
+                    self.is_first = False
+        return segments
+
+    def flush(self) -> str:
+        remaining = self.buffer[self.processed:].strip()
+        if remaining:
+            remaining = _clean_tts_text(remaining)
+        self.processed = len(self.buffer)
+        return remaining
+
+
+# ═══════════════════════════════════════════════════════════════
+#  VAD (dual-threshold + sliding window)
+# ═══════════════════════════════════════════════════════════════
 
 @dataclass
 class _VADState:
@@ -143,7 +210,7 @@ class _VADState:
     silence_counter: int = 0
     total_samples: int = 0
 
-    def reset(self) -> None:
+    def reset(self):
         self.window.clear()
         self.in_speech = False
         self.speech_start_sample = 0
@@ -160,8 +227,7 @@ def _vad_process(state: _VADState, frame: np.ndarray, sample_rate: int) -> bool:
     else:
         is_voice = state.window[-1] if state.window else False
     state.window.append(is_voice)
-    voice_ratio = sum(state.window) / max(len(state.window), 1)
-    have_voice = voice_ratio >= VAD_VOICE_RATIO
+    have_voice = sum(state.window) / max(len(state.window), 1) >= VAD_VOICE_RATIO
 
     if have_voice and not state.in_speech:
         logger.info("VAD START")
@@ -184,7 +250,9 @@ def _vad_min_speech_met(state: _VADState, sample_rate: int) -> bool:
     return (state.total_samples - state.speech_start_sample) / sample_rate >= MIN_SPEECH_S
 
 
-# ── Worker ──
+# ═══════════════════════════════════════════════════════════════
+#  Worker
+# ═══════════════════════════════════════════════════════════════
 
 class Worker:
     def __init__(self, room, identity="asr-bot"):
@@ -218,7 +286,7 @@ class Worker:
         s = p.sid
         if s in self.tasks:
             return
-        logger.info("TRACK %s k=%s", p.identity, t.kind)
+        logger.info("TRACK %s", p.identity)
         self.tasks[s] = asyncio.create_task(self._run(s, t))
 
     def _on_off(self, p):
@@ -240,8 +308,6 @@ class Worker:
             fc = 0
             async for ev in stream:
                 fc += 1
-                if fc == 1:
-                    logger.info("[%s] first audio frame", sid)
                 if busy:
                     continue
                 arr = np.frombuffer(ev.frame.data, dtype=np.int16).astype(np.float32) / 32768.0
@@ -274,12 +340,12 @@ class Worker:
                     try:
                         await self._transcribe_and_respond(buf)
                     except Exception:
-                        logger.exception("[%s] max transcribe failed", sid)
+                        logger.exception("[%s] max failed", sid)
                     busy = False
                     vad.reset()
                     vad.total_samples = len(buf.buffer)
 
-            logger.info("[%s] stream ended after %d frames", sid, fc)
+            logger.info("[%s] stream end fc=%d", sid, fc)
 
         try:
             await _acc()
@@ -303,59 +369,39 @@ class Worker:
         logger.info("ASR: %s", txt[:80])
         await self._send({"type": "utterance", "text": txt, "seq": 0})
 
-        # ── Streaming LLM + incremental TTS ──
+        # ── Streaming LLM + xiaozhi-style TTS ──
         reply_full = ""
-        tts_pending = ""
+        seg = _TTSSegmenter()
         tts_task = None
 
-        async def _play_segment(text: str):
-            try:
-                a = await _tts(text)
-                if a is not None and len(a) > 0:
-                    await self._play(a, 24000)
-            except Exception:
-                logger.exception("TTS segment failed")
-
         try:
             async for delta, is_final in _llm_stream(txt, self.hist):
                 if delta:
                     reply_full += delta
-                    tts_pending += delta
-                    # Send partial reply to client for display
                     await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})
-
-                    # Split at punctuation for incremental TTS
-                    while True:
-                        m = TTS_CHUNK_PATTERN.search(tts_pending)
-                        if not m:
-                            break
-                        pos = m.end()
-                        segment = tts_pending[:pos].strip()
-                        tts_pending = tts_pending[pos:].lstrip()
-                        if segment and len(segment) >= 2:
-                            logger.info("TTS seg: %s", segment[:40])
-                            # Play previous segment if still running
-                            if tts_task:
-                                await tts_task
-                            tts_task = asyncio.create_task(_play_segment(segment))
+                    for s in seg.feed(delta):
+                        if tts_task:
+                            await tts_task
+                        logger.info("TTS seg(%d): %s", len(s), s[:40])
+                        tts_task = asyncio.create_task(_tts(s))
 
                 if is_final:
-                    # Flush remaining text
-                    if tts_pending.strip():
+                    remaining = seg.flush()
+                    if remaining:
                         if tts_task:
                             await tts_task
-                        tts_task = asyncio.create_task(_play_segment(tts_pending.strip()))
+                        tts_task = asyncio.create_task(_tts(remaining))
                     break
-
         except Exception:
             logger.exception("LLM stream failed")
             reply_full = "抱歉,我暂时无法回答。"
 
-        # Wait for final TTS segment to finish
         if tts_task:
             try:
-                await asyncio.wait_for(tts_task, timeout=15)
-            except asyncio.TimeoutError:
+                wav = await asyncio.wait_for(tts_task, timeout=15)
+                if wav is not None and len(wav) > 0:
+                    await self._play(wav, 24000)
+            except (asyncio.TimeoutError, Exception):
                 pass
 
         if reply_full:
@@ -386,7 +432,7 @@ class Worker:
                 reliable=True, topic="transcription",
             )
         except Exception:
-            logger.exception("send failed")
+            pass
 
 
 async def main():