test/app/services/speaker_registry.py

251 lines
9.1 KiB
Python

# -*- 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