# -*- coding: utf-8 -*- """Speaker embedding extraction, registration, and database identification.""" from __future__ import annotations import asyncio import logging import os import tempfile import threading from typing import Any, Optional import librosa import numpy as np import soundfile as sf import torch from app.core.config import settings from app.core.database import pg_speaker_db logger = logging.getLogger(__name__) def _normalize_embedding(embedding: np.ndarray) -> np.ndarray: array = np.asarray(embedding, dtype=np.float32).flatten() norm = float(np.linalg.norm(array)) if norm < 1e-12: return array return array / norm class SpeakerRegistryService: """Owns speaker embedding models and pgvector lookups.""" def __init__(self) -> None: self._pipelines: dict[tuple[str, str], Any] = {} self._lock = threading.Lock() self._inference_lock = threading.BoundedSemaphore(1) def _device(self) -> str: from app.core.device import detect_device return detect_device(settings.DEVICE) def _get_pipeline( self, model_id: Optional[str] = None, model_revision: Optional[str] = None, ) -> Any: effective_model_id = model_id or settings.SV_MODEL effective_revision = model_revision if model_revision is not None else settings.SV_MODEL_REVISION cache_key = (effective_model_id, effective_revision or "") cached = self._pipelines.get(cache_key) if cached is not None: return cached with self._lock: cached = self._pipelines.get(cache_key) if cached is not None: return cached from modelscope.pipelines import pipeline from modelscope.utils.constant import Tasks from app.infrastructure.model_utils import resolve_model_path model_path = resolve_model_path(effective_model_id) device = self._device() kwargs: dict[str, Any] = { "task": Tasks.speaker_verification, "model": model_path, "device": device, } if effective_revision: kwargs["model_revision"] = effective_revision logger.info("正在加载声纹识别模型: %s, device=%s", model_path, device) pipeline_instance = pipeline(**kwargs) if hasattr(pipeline_instance, "device_name"): pipeline_instance.device_name = device model = getattr(pipeline_instance, "model", None) if model is not None and hasattr(model, "to"): pipeline_instance.model = model.to(device) logger.info("声纹识别模型加载成功") self._pipelines[cache_key] = pipeline_instance return pipeline_instance def ensure_loaded(self) -> None: """Eagerly initialize the speaker verification pipeline at startup.""" self._get_pipeline() def extract_embedding_from_audio( self, audio_data: np.ndarray, sample_rate: int = 16000, *, model_id: Optional[str] = None, model_revision: Optional[str] = None, ) -> np.ndarray: audio = np.asarray(audio_data, dtype=np.float32).flatten() if audio.size == 0: raise ValueError("empty audio") if sample_rate != 16000: audio = librosa.resample(audio, orig_sr=sample_rate, target_sr=16000) pipeline_instance = self._get_pipeline(model_id=model_id, model_revision=model_revision) model = getattr(pipeline_instance, "model", None) if model is None: raise RuntimeError("speaker verification model is not initialized") device = getattr(pipeline_instance, "device_name", self._device()) with self._inference_lock: with torch.no_grad(): embeddings = model(torch.as_tensor(audio[None, :]).to(device)) if isinstance(embeddings, torch.Tensor): embedding = embeddings.detach().cpu().numpy()[0] else: embedding = np.asarray(embeddings, dtype=np.float32)[0] return _normalize_embedding(embedding) def extract_embedding_from_file(self, file_path: str) -> np.ndarray: audio_data, sample_rate = librosa.load(file_path, sr=16000, mono=True) return self.extract_embedding_from_audio(audio_data, int(sample_rate)) async def register_file( self, *, name: str, file_path: str, user_id: Optional[str] = None, ) -> dict[str, Optional[str]]: if not pg_speaker_db.is_connected: raise RuntimeError("Speaker database is not connected") loop = asyncio.get_running_loop() embedding = await loop.run_in_executor( None, self.extract_embedding_from_file, file_path, ) speaker = await pg_speaker_db.save_speaker(name, embedding, user_id=user_id) return { "speaker_id": speaker.get("id"), "name": speaker.get("name") or name, "user_id": speaker.get("user_id"), "speaker_model": "CampPlus", } async def identify_embedding( self, embedding: np.ndarray, threshold: Optional[float] = None, ) -> dict[str, Optional[str]]: if not pg_speaker_db.is_connected: return {"speaker_id": None, "name": None, "user_id": None} speaker = await pg_speaker_db.identify_speaker( _normalize_embedding(embedding), threshold if threshold is not None else settings.SV_THRESHOLD, ) return { "speaker_id": speaker.get("id"), "name": speaker.get("name"), "user_id": speaker.get("user_id"), "speaker_model": "CampPlus", } async def identify_file( self, *, file_path: str, threshold: Optional[float] = None, ) -> dict[str, Optional[str]]: loop = asyncio.get_running_loop() embedding = await loop.run_in_executor( None, self.extract_embedding_from_file, file_path, ) return await self.identify_embedding(embedding, threshold=threshold) async def apply_registered_speakers( self, result: Any, threshold: Optional[float] = None, ) -> Any: if not pg_speaker_db.is_connected: return result cache: dict[str, dict[str, Optional[str]]] = {} for segment in getattr(result, "segments", []) or []: embedding = getattr(segment, "speaker_embedding", None) local_speaker_id = getattr(segment, "speaker_id", None) if embedding is None or not local_speaker_id: continue if local_speaker_id not in cache: matched_candidate = await self.identify_embedding( np.asarray(embedding, dtype=np.float32), threshold=threshold, ) if matched_candidate.get("name"): cache[local_speaker_id] = matched_candidate else: continue matched = cache[local_speaker_id] if matched.get("name"): segment.speaker_id = matched.get("speaker_id") or segment.speaker_id segment.speaker_name = matched.get("name") segment.user_id = matched.get("user_id") return result async def list_speakers(self) -> list[dict[str, Optional[str]]]: return await pg_speaker_db.list_speakers() async def delete_speaker(self, speaker_id: str) -> bool: if not speaker_id.isdigit(): return False return await pg_speaker_db.delete_speaker(int(speaker_id)) @staticmethod async def save_upload_to_temp(content: bytes, suffix: str = ".wav") -> str: fd, path = tempfile.mkstemp(prefix="speaker_", suffix=suffix, dir=settings.TEMP_DIR) try: with os.fdopen(fd, "wb") as file_obj: file_obj.write(content) except Exception: os.close(fd) raise return path @staticmethod def cleanup_file(path: Optional[str]) -> None: if path and os.path.exists(path): try: os.remove(path) except Exception as exc: logger.warning("清理临时声纹文件失败 %s: %s", path, exc) @staticmethod def save_audio_array_to_temp(audio_data: np.ndarray, sample_rate: int = 16000) -> str: fd, path = tempfile.mkstemp(prefix="speaker_array_", suffix=".wav", dir=settings.TEMP_DIR) os.close(fd) sf.write(path, np.asarray(audio_data, dtype=np.float32), sample_rate) return path _speaker_registry_service: Optional[SpeakerRegistryService] = None _speaker_registry_lock = threading.Lock() def get_speaker_registry_service() -> SpeakerRegistryService: global _speaker_registry_service if _speaker_registry_service is None: with _speaker_registry_lock: if _speaker_registry_service is None: _speaker_registry_service = SpeakerRegistryService() return _speaker_registry_service