test/app/services/realtime_speaker_clusterer.py

641 lines
25 KiB
Python

# -*- coding: utf-8 -*-
"""FunASR-style realtime speaker chunking and clustering."""
from __future__ import annotations
import logging
import re
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import scipy.linalg
import sklearn.metrics.pairwise
from sklearn.cluster._kmeans import k_means
from app.core.config import settings
from app.core.executor import run_sync
from app.services.speaker_registry import get_speaker_registry_service
logger = logging.getLogger(__name__)
class RealtimeSpeakerClusterer:
"""Sentence-level speaker attribution using chunk embeddings and session clustering."""
FUNASR_SMOOTH_MINDUR_SEC = 0.7
_FAST_ATTACH_MAX_SEC = 1.6
_REGISTRY_MATCH_MIN_SEC = 2.4
def _is_generic_speaker_name(self, value: Any) -> bool:
text = str(value or "").strip()
return bool(re.fullmatch(r"Speaker\d+", text))
def _has_named_identity(self, record: Optional[Dict[str, Any]]) -> bool:
if not record:
return False
return bool(
record.get("registry_speaker_id")
or record.get("user_id")
or (
record.get("speaker_name")
and not self._is_generic_speaker_name(record.get("speaker_name"))
)
)
def _normalize_embedding(self, embedding: np.ndarray) -> np.ndarray:
emb = np.asarray(embedding, dtype=np.float32).reshape(-1)
norm = max(float(np.linalg.norm(emb)), 1e-12)
return (emb / norm).astype(np.float32)
@staticmethod
def _coerce_cluster_index(value: Any, fallback: int) -> int:
try:
if value is None:
return int(fallback)
return int(value)
except (TypeError, ValueError):
return int(fallback)
def _match_existing_speaker(
self,
speaker_records: List[Dict[str, Any]],
mean_embedding: np.ndarray,
) -> Optional[Dict[str, Any]]:
if not speaker_records:
return None
best_record: Optional[Dict[str, Any]] = None
best_score = -1.0
for record in speaker_records:
record_embedding = record.get("embedding")
if record_embedding is None:
continue
score = float(np.dot(self._normalize_embedding(record_embedding), mean_embedding))
if score > best_score:
best_score = score
best_record = record
if best_record is None:
return None
base_threshold = max(
float(getattr(settings, "REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD", 0.58) or 0.58),
0.3,
)
has_named_identity = bool(
best_record.get("registry_speaker_id")
or best_record.get("user_id")
or (
best_record.get("speaker_name")
and not self._is_generic_speaker_name(best_record.get("speaker_name"))
)
)
threshold = max(base_threshold, 0.62 if has_named_identity else 0.72)
if best_score < threshold:
return None
matched = dict(best_record)
matched["_match_score"] = best_score
return matched
def _next_generic_speaker_id(
self,
speaker_records: List[Dict[str, Any]],
) -> str:
seen: set[int] = set()
for record in speaker_records:
for value in (
record.get("speaker_id"),
record.get("speaker_name"),
):
text = str(value or "").strip()
match = re.fullmatch(r"Speaker(\d+)", text)
if match:
seen.add(int(match.group(1)))
next_index = max(seen, default=0) + 1
return f"Speaker{next_index:02d}"
def _is_mixed_speaker_segment(
self,
speaker_records: List[Dict[str, Any]],
current_chunks: List[Dict[str, Any]],
) -> bool:
if len(current_chunks) < 2 or len(speaker_records) < 2:
return False
assigned_labels: List[str] = []
for chunk in current_chunks:
embedding = chunk.get("embedding")
if embedding is None:
continue
matched = self._match_existing_speaker(
speaker_records,
self._normalize_embedding(np.asarray(embedding, dtype=np.float32)),
)
if not matched:
continue
label = str(
matched.get("registry_speaker_id")
or matched.get("speaker_name")
or matched.get("speaker_id")
or ""
).strip()
if label:
assigned_labels.append(label)
if len(assigned_labels) < 2:
return False
return len(set(assigned_labels)) >= 2
def _correct_labels(self, labels: np.ndarray) -> np.ndarray:
labels_id = 0
id2id: dict[int, int] = {}
new_labels: list[int] = []
for label in labels.tolist():
label = int(label)
if label not in id2id:
id2id[label] = labels_id
labels_id += 1
new_labels.append(id2id[label])
return np.asarray(new_labels, dtype=np.int32)
def _spectral_cluster(self, X: np.ndarray, oracle_num: Optional[int] = None) -> np.ndarray:
sim_mat = sklearn.metrics.pairwise.cosine_similarity(X, X)
A = sim_mat.copy()
pval = 0.022
if A.shape[0] * pval < 6:
pval = 6.0 / A.shape[0]
n_elems = int((1 - pval) * A.shape[0])
for i in range(A.shape[0]):
low_indexes = np.argsort(A[i, :])[:n_elems]
A[i, low_indexes] = 0
A = 0.5 * (A + A.T)
A[np.diag_indices(A.shape[0])] = 0
D = np.diag(np.sum(np.abs(A), axis=1))
L = D - A
lambdas, eig_vecs = scipy.linalg.eigh(L)
if oracle_num is not None:
num_spk = max(int(oracle_num), 1)
else:
min_num_spks = 1
max_num_spks = min(15, max(1, X.shape[0] - 1))
eig_slice = lambdas[min_num_spks - 1 : max_num_spks + 1]
gap_list = [float(eig_slice[i + 1]) - float(eig_slice[i]) for i in range(len(eig_slice) - 1)]
num_spk = int(np.argmax(gap_list)) + min_num_spks if gap_list else 1
emb = eig_vecs[:, : max(num_spk, 1)]
_, labels, _ = k_means(emb, max(num_spk, 1))
return np.asarray(labels, dtype=np.int32)
def _merge_by_cos(self, labels: np.ndarray, embs: np.ndarray, cos_thr: float) -> np.ndarray:
labels = np.asarray(labels, dtype=np.int32).copy()
while True:
spk_num = int(labels.max()) + 1
if spk_num <= 1:
break
centers = []
for i in range(spk_num):
spk_emb = embs[labels == i].mean(0)
centers.append(spk_emb)
centers = np.stack(centers, axis=0)
norm_centers = centers / np.linalg.norm(centers, axis=1, keepdims=True)
affinity = np.matmul(norm_centers, norm_centers.T)
affinity = np.triu(affinity, 1)
spks = np.unravel_index(np.argmax(affinity), affinity.shape)
if float(affinity[spks]) < cos_thr:
break
for i in range(len(labels)):
if labels[i] == spks[1]:
labels[i] = spks[0]
elif labels[i] > spks[1]:
labels[i] -= 1
return self._correct_labels(labels)
def _cluster_embeddings(self, X: np.ndarray, oracle_num: Optional[int] = None) -> np.ndarray:
if X.shape[0] < 20:
labels = np.zeros(X.shape[0], dtype=np.int32)
else:
labels = self._spectral_cluster(X, oracle_num=oracle_num)
return self._merge_by_cos(labels, X, cos_thr=0.78 if oracle_num is None else 1.0)
def build_sv_chunks(
self,
audio: np.ndarray,
*,
segment_start_ms: int = 0,
) -> List[Dict[str, Any]]:
audio = np.asarray(audio, dtype=np.float32).flatten()
if audio.size == 0:
return []
sample_rate = 16000
chunk_len = int(1.5 * sample_rate)
chunk_shift = int(0.75 * sample_rate)
chunks: List[Dict[str, Any]] = []
last_chunk_end = 0
for chunk_start in range(0, audio.shape[0], chunk_shift):
chunk_end = min(chunk_start + chunk_len, audio.shape[0])
if chunk_end <= last_chunk_end:
break
actual_start = max(0, chunk_end - chunk_len)
actual_end = chunk_end
last_chunk_end = actual_end
chunk_audio = np.asarray(audio[actual_start:actual_end], dtype=np.float32)
if chunk_audio.shape[0] < chunk_len:
chunk_audio = np.pad(chunk_audio, (0, chunk_len - chunk_audio.shape[0]), "constant")
chunks.append(
{
"start_ms": int(segment_start_ms + actual_start / 16.0),
"end_ms": int(segment_start_ms + actual_end / 16.0),
"audio": chunk_audio,
}
)
return chunks
async def extract_chunk_embeddings(
self,
audio: np.ndarray,
*,
segment_start_ms: int = 0,
) -> List[Dict[str, Any]]:
chunk_items = self.build_sv_chunks(audio, segment_start_ms=segment_start_ms)
if not chunk_items:
chunk_items = [
{
"start_ms": int(segment_start_ms),
"end_ms": int(segment_start_ms + np.asarray(audio, dtype=np.float32).size / 16.0),
"audio": np.asarray(audio, dtype=np.float32),
}
]
registry = get_speaker_registry_service()
results: List[Dict[str, Any]] = []
for item in chunk_items:
embedding = await run_sync(
registry.extract_embedding_from_audio,
item["audio"],
16000,
model_id=settings.REALTIME_SV_MODEL,
model_revision=settings.REALTIME_SV_MODEL_REVISION or None,
)
results.append(
{
"start_ms": int(item["start_ms"]),
"end_ms": int(item["end_ms"]),
"embedding": self._normalize_embedding(np.asarray(embedding, dtype=np.float32)),
}
)
return results
async def extract_registry_embedding(self, audio: np.ndarray) -> Optional[np.ndarray]:
audio_array = np.asarray(audio, dtype=np.float32).flatten()
if audio_array.size == 0:
return None
try:
registry = get_speaker_registry_service()
embedding = await run_sync(
registry.extract_embedding_from_audio,
audio_array,
16000,
)
return self._normalize_embedding(np.asarray(embedding, dtype=np.float32))
except Exception as exc:
logger.warning("Realtime registry speaker embedding failed: %s", exc)
return None
def _flatten_chunk_records(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], np.ndarray]:
flat_chunks: List[Dict[str, Any]] = []
flat_embeddings: List[np.ndarray] = []
for record_idx, record in enumerate(records):
chunk_entries = record.get("chunks") or []
if not chunk_entries:
embeddings = record.get("embeddings") or []
if not embeddings and record.get("embedding") is not None:
embeddings = [record["embedding"]]
start_ms = int(record.get("start_ms", 0))
end_ms = int(record.get("end_ms", start_ms))
for embedding in embeddings:
chunk_entries.append(
{
"start_ms": start_ms,
"end_ms": end_ms,
"embedding": embedding,
}
)
for chunk in chunk_entries:
emb = self._normalize_embedding(np.asarray(chunk["embedding"], dtype=np.float32))
flat_chunks.append(
{
"record_index": record_idx,
"start_ms": int(chunk.get("start_ms", record.get("start_ms", 0))),
"end_ms": int(chunk.get("end_ms", record.get("end_ms", 0))),
}
)
flat_embeddings.append(emb)
if not flat_embeddings:
return [], np.zeros((0, 0), dtype=np.float32)
return flat_chunks, np.stack(flat_embeddings, axis=0).astype(np.float32)
def _merge_seque(self, distribute_res: List[List[float]]) -> List[List[float]]:
if not distribute_res:
return []
res = [distribute_res[0][:]]
for item in distribute_res[1:]:
if item[2] != res[-1][2] or item[0] > res[-1][1]:
res.append(item[:])
else:
res[-1][1] = item[1]
return res
def _smooth_timeline(
self,
res: List[List[float]],
mindur: float = FUNASR_SMOOTH_MINDUR_SEC,
) -> List[List[float]]:
if len(res) < 2:
return res
for item in res:
item[0] = round(float(item[0]), 2)
item[1] = round(float(item[1]), 2)
for idx in range(len(res)):
if res[idx][1] - res[idx][0] < mindur:
if idx == 0:
res[idx][2] = res[idx + 1][2]
elif idx == len(res) - 1:
res[idx][2] = res[idx - 1][2]
elif res[idx][0] - res[idx - 1][1] <= res[idx + 1][0] - res[idx][1]:
res[idx][2] = res[idx - 1][2]
else:
res[idx][2] = res[idx + 1][2]
return self._merge_seque(res)
def _postprocess_timeline(
self,
flat_chunks: List[Dict[str, Any]],
labels: np.ndarray,
embeddings: np.ndarray,
) -> List[Dict[str, Any]]:
assert len(flat_chunks) == len(labels)
labels = self._correct_labels(labels)
distribute_res = [
[
float(chunk["start_ms"]) / 1000.0,
float(chunk["end_ms"]) / 1000.0,
int(labels[idx]),
]
for idx, chunk in enumerate(flat_chunks)
]
distribute_res = self._merge_seque(distribute_res)
def is_overlapped(t1: float, t2: float) -> bool:
return t1 > t2 + 1e-4
for idx in range(1, len(distribute_res)):
if is_overlapped(distribute_res[idx - 1][1], distribute_res[idx][0]):
pivot = (distribute_res[idx][0] + distribute_res[idx - 1][1]) / 2.0
distribute_res[idx][0] = pivot
distribute_res[idx - 1][1] = pivot
distribute_res = self._smooth_timeline(distribute_res)
return [
{
"start_ms": int(round(item[0] * 1000.0)),
"end_ms": int(round(item[1] * 1000.0)),
"cluster_index": int(item[2]),
}
for item in distribute_res
]
def _pick_segment_cluster(
self,
cluster_ranges: List[Dict[str, Any]],
*,
segment_start_ms: int,
segment_end_ms: int,
) -> int:
overlaps: Dict[int, int] = defaultdict(int)
for item in cluster_ranges:
overlap = min(segment_end_ms, int(item["end_ms"])) - max(segment_start_ms, int(item["start_ms"]))
if overlap > 0:
overlaps[int(item["cluster_index"])] += int(overlap)
if overlaps:
return max(overlaps.items(), key=lambda kv: (kv[1], -kv[0]))[0]
centers = [
item for item in cluster_ranges
if int(item["start_ms"]) <= segment_end_ms and int(item["end_ms"]) >= segment_start_ms
]
if centers:
return int(centers[0]["cluster_index"])
return 0
def cluster_records_with_ranges(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], List[int], List[Dict[str, Any]]]:
flat_chunks, X = self._flatten_chunk_records(records)
if X.size == 0:
return [], [], []
labels = self._cluster_embeddings(X, oracle_num=None)
cluster_ranges = self._postprocess_timeline(flat_chunks, labels, X)
cluster_ids = sorted({int(item["cluster_index"]) for item in cluster_ranges})
if not cluster_ids:
cluster_ids = sorted({int(label) for label in labels.tolist()})
clusters: List[Dict[str, Any]] = []
for cluster_idx in cluster_ids:
member_mask = labels == int(cluster_idx)
member_embeddings = X[member_mask]
centroid = self._normalize_embedding(member_embeddings.mean(0))
record_refs = []
seen_record_indices = set()
for emb_idx, chunk in enumerate(flat_chunks):
record_idx = int(chunk["record_index"])
if not member_mask[emb_idx] or record_idx in seen_record_indices:
continue
seen_record_indices.add(record_idx)
record_refs.append(records[record_idx])
clusters.append(
{
"centroid": centroid,
"count": int(member_mask.sum()),
"record_refs": record_refs,
}
)
record_to_labels: List[List[int]] = [[] for _ in records]
for emb_idx, chunk in enumerate(flat_chunks):
record_to_labels[int(chunk["record_index"])].append(int(labels[emb_idx]))
record_assignments: List[int] = []
for record, local_assignments in zip(records, record_to_labels):
if local_assignments:
cluster_idx = self._pick_segment_cluster(
cluster_ranges,
segment_start_ms=int(record.get("start_ms", 0)),
segment_end_ms=int(record.get("end_ms", record.get("start_ms", 0))),
)
record_assignments.append(cluster_idx)
else:
record_assignments.append(-1)
return clusters, record_assignments, cluster_ranges
def cluster_records(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], List[int]]:
clusters, record_assignments, _ = self.cluster_records_with_ranges(records)
return clusters, record_assignments
async def resolve_segment_speaker(
self,
speaker_records: List[Dict[str, Any]],
audio: np.ndarray,
*,
segment_start_ms: int,
segment_end_ms: int,
enable_registry_match: bool,
speaker_threshold: Optional[float],
) -> Optional[Dict[str, Any]]:
duration_sec = float(len(audio)) / 16000.0
last_record = speaker_records[-1] if speaker_records else None
if duration_sec < self._FAST_ATTACH_MAX_SEC and speaker_records:
if self._has_named_identity(last_record):
speaker_id = last_record.get("registry_speaker_id") or last_record.get("speaker_id") or "Speaker01"
speaker_name = last_record.get("speaker_name") or speaker_id
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": last_record.get("user_id"),
"registry_speaker_id": last_record.get("registry_speaker_id"),
"speaker_confidence": 0.0,
"speaker_strategy": "short_attach",
"_cluster_index": last_record.get("cluster_index"),
"_embedding": last_record.get("embedding"),
"_chunk_embeddings": last_record.get("embeddings") or [],
"_chunks": last_record.get("chunks") or [],
}
current_chunks = await self.extract_chunk_embeddings(
np.asarray(audio, dtype=np.float32),
segment_start_ms=segment_start_ms,
)
if not current_chunks:
if speaker_records:
if self._has_named_identity(last_record):
speaker_id = last_record.get("registry_speaker_id") or last_record.get("speaker_id") or "Speaker01"
speaker_name = last_record.get("speaker_name") or speaker_id
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": last_record.get("user_id"),
"registry_speaker_id": last_record.get("registry_speaker_id"),
"speaker_confidence": 0.0,
"speaker_strategy": "embedding_attach",
"_cluster_index": last_record.get("cluster_index"),
"_embedding": last_record.get("embedding"),
"_chunk_embeddings": last_record.get("embeddings") or [],
"_chunks": last_record.get("chunks") or [],
}
return None
mean_embedding = self._normalize_embedding(
np.mean(np.stack([chunk["embedding"] for chunk in current_chunks], axis=0), axis=0)
)
if duration_sec >= 4.0 and self._is_mixed_speaker_segment(speaker_records, current_chunks):
return {
"speaker_id": -1,
"speaker_name": "",
"user_id": None,
"registry_speaker_id": None,
"speaker_confidence": 0.0,
"speaker_strategy": "mixed_segment",
"_cluster_index": None,
"_embedding": mean_embedding,
"_chunk_embeddings": [np.asarray(chunk["embedding"], dtype=np.float32) for chunk in current_chunks],
"_chunks": current_chunks,
}
matched_record = self._match_existing_speaker(speaker_records, mean_embedding)
speaker_id = self._next_generic_speaker_id(speaker_records)
speaker_name = speaker_id
user_id = None
registry_speaker_id = None
strategy = "new_speaker"
cluster_index = max(
[
int(record.get("cluster_index", -1))
for record in speaker_records
if record.get("cluster_index") is not None
] or [-1]
) + 1
confidence = 0.0
if matched_record is not None:
speaker_id = (
matched_record.get("registry_speaker_id")
or matched_record.get("speaker_id")
or speaker_id
)
speaker_name = matched_record.get("speaker_name") or speaker_id
user_id = matched_record.get("user_id")
registry_speaker_id = matched_record.get("registry_speaker_id")
cluster_index = self._coerce_cluster_index(
matched_record.get("cluster_index"),
cluster_index,
)
confidence = float(matched_record.get("_match_score", 0.0))
strategy = "embedding_match"
if (
not registry_speaker_id
and enable_registry_match
and duration_sec >= self._REGISTRY_MATCH_MIN_SEC
and (
matched_record is None
or confidence >= max(float(getattr(settings, "REALTIME_SPEAKER_CONFIRM_THRESHOLD", 0.62) or 0.62), 0.7)
)
):
# Realtime clustering can use a dedicated fast model, while the
# registry stores embeddings from the registration model. Match the
# registry in its own embedding space instead of comparing vectors
# extracted by a different model.
registry_embedding = await self.extract_registry_embedding(audio)
if registry_embedding is not None:
matched = await get_speaker_registry_service().identify_embedding(
registry_embedding,
threshold=speaker_threshold,
)
if matched.get("name"):
speaker_id = matched.get("speaker_id") or speaker_id
speaker_name = matched.get("name") or speaker_name
user_id = matched.get("user_id")
registry_speaker_id = matched.get("speaker_id")
strategy = "registry_match"
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": user_id,
"registry_speaker_id": registry_speaker_id,
"speaker_confidence": round(confidence, 4),
"speaker_strategy": strategy,
"_cluster_index": self._coerce_cluster_index(cluster_index, 0),
"_embedding": mean_embedding,
"_chunk_embeddings": [np.asarray(chunk["embedding"], dtype=np.float32) for chunk in current_chunks],
"_chunks": current_chunks,
"_matched_existing": matched_record is not None,
}
_realtime_speaker_clusterer: Optional[RealtimeSpeakerClusterer] = None
def get_realtime_speaker_clusterer() -> RealtimeSpeakerClusterer:
global _realtime_speaker_clusterer
if _realtime_speaker_clusterer is None:
_realtime_speaker_clusterer = RealtimeSpeakerClusterer()
return _realtime_speaker_clusterer