|
|
@@ -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
|