|
|
@@ -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():
|