Bladeren bron

feat(voiceprint): Zvec-backed voiceprint registration + identification

- voiceprint.py: Resemblyzer encoder + Zvec vector store
- Multi-user registration ('注册声纹 张三' / '我是张三')
- Parallel identification with ASR, injects speaker name into LLM
- Zvec embedded DB, no external service needed
- Dockerfile: added resemblyzer + zvec pip install
wenhongquan 3 weken geleden
bovenliggende
commit
4abc018a4f
2 gewijzigde bestanden met toevoegingen van 254 en 1 verwijderingen
  1. 215 0
      asr_agent/voiceprint.py
  2. 39 1
      asr_agent/worker.py

+ 215 - 0
asr_agent/voiceprint.py

@@ -0,0 +1,215 @@
+"""
+Voiceprint (声纹) Recognition — Zvec-backed vector store
+- Multi-user registration, identification, management
+- Resemblyzer (GE2E) for speaker embedding extraction
+- Zvec embedded vector DB for storage & similarity search
+- Runs parallel to ASR on the same audio buffer
+"""
+
+from __future__ import annotations
+
+import logging
+import os
+import time
+from typing import Optional
+
+import numpy as np
+
+logger = logging.getLogger("voiceprint")
+
+VP_DB_PATH = os.environ.get("VP_DB_PATH", "/data/voiceprints")
+VP_SIMILARITY_THRESHOLD = float(os.environ.get("VP_SIMILARITY_THRESHOLD", "0.65"))
+
+# ── Lazy-loaded encoder ──
+_encoder = None
+
+
+def _get_encoder():
+    global _encoder
+    if _encoder is not None:
+        return _encoder
+    try:
+        from resemblyzer import VoiceEncoder
+        _encoder = VoiceEncoder(device="cpu", verbose=False)
+        logger.info("Voiceprint encoder loaded")
+    except ImportError:
+        logger.warning("resemblyzer not installed — voiceprint disabled")
+        return None
+    return _encoder
+
+
+# ── Zvec store ──
+_collection = None
+_EMBEDDING_DIM = 256  # Resemblyzer GE2E output dimension
+
+
+def _get_collection():
+    global _collection
+    if _collection is not None:
+        return _collection
+    try:
+        import zvec
+
+        schema = zvec.CollectionSchema(
+            name="voiceprints",
+            fields=[
+                zvec.FieldSchema(name="display_name", data_type=zvec.DataType.STRING),
+                zvec.FieldSchema(name="registered_at", data_type=zvec.DataType.FLOAT64),
+                zvec.FieldSchema(name="last_matched_at", data_type=zvec.DataType.FLOAT64),
+            ],
+            vectors=[
+                zvec.VectorSchema(
+                    name="embedding",
+                    data_type=zvec.DataType.VECTOR_FP32,
+                    dimension=_EMBEDDING_DIM,
+                    index_param=zvec.HnswIndexParam(metric_type=zvec.MetricType.COSINE),
+                ),
+            ],
+        )
+
+        os.makedirs(VP_DB_PATH, exist_ok=True)
+        _collection = zvec.create_and_open(path=VP_DB_PATH, schema=schema)
+        _collection.optimize()
+        logger.info("Zvec voiceprint DB opened at %s (%d voices)", VP_DB_PATH, _collection.stats.row_count)
+    except ImportError:
+        logger.warning("zvec not installed — voiceprint disabled")
+        return None
+    except Exception:
+        logger.exception("Failed to open Zvec voiceprint DB")
+        return None
+    return _collection
+
+
+class VoiceprintStore:
+    """Zvec-backed voiceprint database."""
+
+    def __init__(self):
+        self._col = _get_collection()
+
+    @property
+    def enabled(self) -> bool:
+        return self._col is not None
+
+    # ── Registration ──
+
+    def register(self, user_id: str, display_name: str, audio: np.ndarray) -> bool:
+        """Register a voiceprint from audio (float32, 16kHz mono). Returns True on success."""
+        enc = _get_encoder()
+        col = self._col
+        if enc is None or col is None:
+            return False
+        if len(audio) < 8000:
+            logger.warning("Audio too short for voiceprint (%d samples)", len(audio))
+            return False
+
+        emb = enc.embed_utterance(audio)
+        if emb is None:
+            logger.warning("Voiceprint extraction failed")
+            return False
+
+        import zvec
+        # Delete existing voiceprint for this user if any
+        try:
+            col.delete(ids=user_id)
+        except Exception:
+            pass
+
+        col.insert(zvec.Doc(
+            id=user_id,
+            vectors={"embedding": emb.tolist()},
+            fields={
+                "display_name": display_name,
+                "registered_at": time.time(),
+                "last_matched_at": 0.0,
+            },
+        ))
+        col.optimize()
+        logger.info("Voiceprint registered: %s (%s)", display_name, user_id)
+        return True
+
+    # ── Identification ──
+
+    def identify(self, audio: np.ndarray) -> tuple[Optional[str], Optional[str], float]:
+        """Identify speaker. Returns (user_id, display_name, similarity)."""
+        enc = _get_encoder()
+        col = self._col
+        if enc is None or col is None:
+            return None, None, 0.0
+        if len(audio) < 8000:
+            return None, None, 0.0
+
+        emb = enc.embed_utterance(audio)
+        if emb is None:
+            return None, None, 0.0
+
+        import zvec
+        result = col.query(
+            queries=zvec.Query(field_name="embedding", vector=emb.tolist()),
+            topk=1,
+        )
+        if not result:
+            return None, None, 0.0
+
+        doc = result[0]
+        sim = float(doc.score)
+        if sim >= VP_SIMILARITY_THRESHOLD:
+            uid = doc.id
+            name = doc.fields.get("display_name", uid)
+            # Update last matched time
+            try:
+                col.insert(zvec.Doc(
+                    id=uid,
+                    vectors={"embedding": emb.tolist()},  # refresh embedding
+                    fields={
+                        "display_name": name,
+                        "registered_at": doc.fields.get("registered_at", time.time()),
+                        "last_matched_at": time.time(),
+                    },
+                ))
+                col.optimize()
+            except Exception:
+                pass
+            return uid, name, sim
+        return None, None, sim
+
+    # ── Management ──
+
+    def list_users(self) -> list[dict]:
+        """List all registered voiceprints."""
+        col = self._col
+        if col is None:
+            return []
+        results = []
+        stats = col.stats
+        # Zvec doesn't have a "list all" directly; we fetch by known IDs
+        # For now, return count only
+        return [{"count": stats.row_count}]
+
+    def delete(self, user_id: str) -> bool:
+        """Delete a voiceprint by user_id."""
+        col = self._col
+        if col is None:
+            return False
+        try:
+            col.delete(ids=user_id)
+            col.optimize()
+            logger.info("Voiceprint deleted: %s", user_id)
+            return True
+        except Exception:
+            return False
+
+    def get_speaker_context(self, speaker_name: str | None) -> str:
+        if speaker_name:
+            return f"当前说话人是 {speaker_name},请用对方习惯的方式回应。"
+        return ""
+
+
+# ── Singleton ──
+_store: VoiceprintStore | None = None
+
+
+def get_store() -> VoiceprintStore:
+    global _store
+    if _store is None:
+        _store = VoiceprintStore()
+    return _store

+ 39 - 1
asr_agent/worker.py

@@ -36,6 +36,14 @@ sys.path.insert(0, _root)
 from audio_processor import AudioBuffer
 from qwen_engine import QwenASREngine
 
+try:
+    from voiceprint import get_store
+    _VP_ENABLED = True
+except ImportError:
+    _VP_ENABLED = False
+    def get_store():
+        return None
+
 logger = logging.getLogger("worker")
 
 # ── Env ──
@@ -615,13 +623,43 @@ class Worker:
         logger.info("ASR: %s", txt[:80])
         await self._send({"type": "utterance", "text": txt, "seq": 0})
 
+        # ── Voiceprint identification (parallel with LLM prep) ──
+        speaker_name = None
+        if _VP_ENABLED:
+            store = get_store()
+            if store and store.enabled:
+                loop = asyncio.get_running_loop()
+                _id, speaker_name, _sim = await loop.run_in_executor(None, store.identify, all_audio)
+                if speaker_name:
+                    logger.info("VP: identified %s (sim=%.3f)", speaker_name, _sim)
+                    await self._send({"type": "speaker", "name": speaker_name, "confidence": round(_sim, 3)})
+
+        # ── Voiceprint registration command ──
+        import re
+        vp_match = re.match(r"(注册声纹|我是|我叫)\s*(.+)", txt)
+        if vp_match and _VP_ENABLED:
+            store = get_store()
+            if store and store.enabled:
+                name = vp_match.group(2).strip()
+                user_id = f"user_{hash(name) % 1000000:06d}"
+                ok = await loop.run_in_executor(None, store.register, user_id, name, all_audio)
+                if ok:
+                    await self._send({"type": "vp_registered", "name": name, "user_id": user_id})
+                    speaker_name = name
+                    logger.info("VP: registered %s", name)
+
         logger.info("LLM start")
         reply_full = ""
         seg = _TTSSegmenter()
         tts_segments: list[str] = []
 
+        # ── Build prompt with voiceprint context ──
+        prompt = txt
+        if speaker_name:
+            prompt = f"[说话人: {speaker_name}] {txt}"
+
         try:
-            async for delta, is_final in _llm_stream(txt, self.hist):
+            async for delta, is_final in _llm_stream(prompt, self.hist):
                 if delta:
                     reply_full += delta
                     await self._send({"type": "reply_partial", "text": reply_full, "seq": 0})