Преглед изворни кода

feat(provider): add LLM/ASR provider switch (vllm/mimo, qwen/mimo)

- LLM_PROVIDER env (vllm|mimo): switches LLM backend
- ASR_PROVIDER env (qwen|mimo): switches ASR backend
- Mimo LLM: streaming via mimo-v2.5 API (same format as vLLM)
- Mimo ASR: mio-v2.5-asr via API (WAV base64 upload)
- Existing vLLM + Qwen3 remain default, no breaking change
wenhongquan пре 3 недеља
родитељ
комит
a381dc1067
1 измењених фајлова са 99 додато и 2 уклоњено
  1. 99 2
      asr_agent/conversation_worker.py

+ 99 - 2
asr_agent/conversation_worker.py

@@ -44,6 +44,11 @@ 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")
 LLM_MODEL = os.environ.get("LLM_MODEL", "qwen3.6-35b-awq")
 MIMO_KEY = os.environ.get("MIMO_KEY", "")
+MIMO_API_BASE = os.environ.get("MIMO_API_BASE", "https://token-plan-cn.xiaomimimo.com/v1")
+
+# Provider switches: "vllm" | "mimo" for LLM, "qwen" | "mimo" for ASR
+LLM_PROVIDER = os.environ.get("LLM_PROVIDER", "vllm")
+ASR_PROVIDER = os.environ.get("ASR_PROVIDER", "qwen")
 
 
 # ── LLM ──
@@ -83,6 +88,17 @@ def _token(room, ident):
 # ═══════════════════════════════════════════════════════════════
 
 async def _llm_stream(prompt: str, hist: list[dict]):
+    """Stream LLM tokens, dispatching based on LLM_PROVIDER env."""
+    if LLM_PROVIDER == "mimo":
+        async for x in _llm_stream_mimo(prompt, hist):
+            yield x
+        return
+    async for x in _llm_stream_vllm(prompt, hist):
+        yield x
+
+
+async def _llm_stream_vllm(prompt: str, hist: list[dict]):
+    """vLLM streaming via OpenAI-compatible SSE."""
     msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
     msgs.extend(hist)
     msgs.append({"role": "user", "content": prompt})
@@ -119,6 +135,88 @@ async def _llm_stream(prompt: str, hist: list[dict]):
     yield "", True
 
 
+async def _llm_stream_mimo(prompt: str, hist: list[dict]):
+    """Mimo v2.5 LLM streaming via OpenAI-compatible SSE."""
+    if not MIMO_KEY:
+        yield "", True
+        return
+    msgs = [{"role": "system", "content": SYSTEM_PROMPT}]
+    msgs.extend(hist)
+    msgs.append({"role": "user", "content": prompt})
+
+    h = {"api-key": MIMO_KEY, "Content-Type": "application/json"}
+    async with aiohttp.ClientSession() as s:
+        async with s.post(
+            f"{MIMO_API_BASE}/chat/completions",
+            json={"model": "mimo-v2.5", "messages": msgs,
+                  "max_tokens": LLM_MAX_TOKENS, "temperature": LLM_TEMPERATURE,
+                  "stream": True},
+            headers=h, timeout=aiohttp.ClientTimeout(total=LLM_TIMEOUT),
+        ) as r:
+            async for line in r.content:
+                line = line.decode().strip()
+                if not line.startswith("data: "):
+                    continue
+                data = line[6:]
+                if data == "[DONE]":
+                    break
+                try:
+                    chunk = json.loads(data)
+                    delta = chunk.get("choices", [{}])[0].get("delta", {})
+                    text = delta.get("content", "")
+                    if text and text.strip():
+                        yield text, False
+                except Exception:
+                    continue
+    yield "", True
+
+
+# ═══════════════════════════════════════════════════════════════
+#  ASR (dispatch: qwen-local | mimo-api)
+# ═══════════════════════════════════════════════════════════════
+
+async def _asr_transcribe(audio_data: np.ndarray, asr_engine=None) -> str:
+    """Dispatch ASR based on ASR_PROVIDER env."""
+    if ASR_PROVIDER == "mimo":
+        return await _asr_mimo(audio_data)
+    # Default: local Qwen3-ASR
+    loop = asyncio.get_running_loop()
+    result = await loop.run_in_executor(None, asr_engine.transcribe_array, audio_data)
+    return result.get("text", "").strip()
+
+
+async def _asr_mimo(audio_data: np.ndarray) -> str:
+    """Mimo v2.5 ASR via API."""
+    if not MIMO_KEY:
+        return ""
+    import io, wave
+    buf = io.BytesIO()
+    with wave.open(buf, "wb") as wf:
+        wf.setnchannels(1)
+        wf.setsampwidth(2)
+        wf.setframerate(16000)
+        wf.writeframes((audio_data * 32767).astype(np.int16).tobytes())
+    audio_b64 = base64.b64encode(buf.getvalue()).decode()
+
+    h = {"api-key": MIMO_KEY, "Content-Type": "application/json"}
+    async with aiohttp.ClientSession() as s:
+        async with s.post(
+            f"{MIMO_API_BASE}/chat/completions",
+            json={"model": "mimo-v2.5-asr", "messages": [
+                {"role": "user", "content": [
+                    {"type": "input_audio", "input_audio": {"data": audio_b64, "format": "wav"}}
+                ]}
+            ]},
+            headers=h, timeout=aiohttp.ClientTimeout(total=20),
+        ) as r:
+            try:
+                result = await r.json()
+                return result.get("choices", [{}])[0].get("message", {}).get("content", "")
+            except Exception:
+                logger.exception("Mimo ASR failed")
+                return ""
+
+
 # ═══════════════════════════════════════════════════════════════
 #  TTS
 # ═══════════════════════════════════════════════════════════════
@@ -433,8 +531,7 @@ class Worker:
             return
 
         loop = asyncio.get_running_loop()
-        result = await loop.run_in_executor(None, self.asr.transcribe_array, all_audio)
-        txt = result.get("text", "").strip()
+        txt = await _asr_transcribe(all_audio, self.asr)
         if not txt:
             return