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