641 lines
25 KiB
Python
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
|