251 lines
9.1 KiB
Python
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
|