diff --git a/.env.example b/.env.example index 7f8b389..5764f9a 100644 --- a/.env.example +++ b/.env.example @@ -18,6 +18,12 @@ FUNASR_SPEAKER_WINDOW_MS=1500 FUNASR_SPEAKER_HOP_MS=750 # CAM++ 每批最多计算多少个窗口,避免很长 turn 一次性占用过多显存。 FUNASR_SPEAKER_BATCH_SIZE=16 +# FunASR HybridSpeakerTracker history and stable identity matching defaults. +FUNASR_SPEAKER_HISTORY_CHUNKS=128 +FUNASR_SPEAKER_MATCH_THRESHOLD=0.6 +FUNASR_SPEAKER_MERGE_THRESHOLD=0.78 +# FunASR's default speaker identity limit; increase this value for larger meetings. +FUNASR_MAX_SPEAKERS=15 # 小于该时长的相邻说话人段会合并,避免短窗噪声把一句话切得过碎。 FUNASR_SPEAKER_MIN_SEGMENT_MS=3000 diff --git a/backend/auxiliary_server.py b/backend/auxiliary_server.py index d14c0f9..8d299e9 100644 --- a/backend/auxiliary_server.py +++ b/backend/auxiliary_server.py @@ -9,6 +9,7 @@ import os import sys import tempfile import time +from collections import deque from collections.abc import Mapping from pathlib import Path from typing import Any @@ -30,7 +31,11 @@ AUXILIARY_PORT = 8010 load_dotenv(PROJECT_ROOT / ".env") MODELS_DIR = Path(os.getenv("MODEL_DIR", str(PROJECT_ROOT / "models"))).resolve() AUXILIARY_DEVICE = os.getenv("AUXILIARY_DEVICE", "cuda:0") -ONLINE_SPEAKER_MATCH_THRESHOLD = 0.68 +# Keep the real-time tracker defaults aligned with FunASR HybridSpeakerTracker. +ONLINE_SPEAKER_MATCH_THRESHOLD = float(os.getenv("FUNASR_SPEAKER_MATCH_THRESHOLD", "0.6")) +ONLINE_CLUSTER_MERGE_THRESHOLD = float(os.getenv("FUNASR_SPEAKER_MERGE_THRESHOLD", "0.78")) +ONLINE_MAX_SPEAKERS = max(1, int(os.getenv("FUNASR_MAX_SPEAKERS", "15"))) +ONLINE_SPEAKER_HISTORY_CHUNKS = max(1, int(os.getenv("FUNASR_SPEAKER_HISTORY_CHUNKS", "128"))) MIN_ONLINE_SPEAKER_AUDIO_MS = 800 # Realtime VAD is loaded by the WebSocket process. This service loads CAM++. # Its standalone VAD HTTP endpoint remains available on demand. @@ -94,7 +99,11 @@ class AuxiliaryRuntime: self.inference_lock = asyncio.Lock() # 每个 WebSocket session 独立维护聚类中心,避免不同浏览器会话互相污染。 self.speaker_clusters: dict[str, list[dict[str, Any]]] = {} + # Keep a bounded rolling window history per WebSocket session, like FunASR's tracker. + self.speaker_history: dict[str, dict[str, Any]] = {} self.speaker_last_seen: dict[str, float] = {} + self._speaker_cluster_backend: Any | None = None + self._speaker_postprocess: Any | None = None def _preload_kinds(self) -> set[str]: """读取需要在启动时加载的模型类型,默认不加载完整 diarization。""" @@ -503,24 +512,124 @@ class AuxiliaryRuntime: for (start_ms, end_ms), embedding in zip(spans, embeddings) ] - @staticmethod - def _speaker_runs(cells: list[dict[str, Any]]) -> list[dict[str, Any]]: - """把重叠声纹窗变成连续时间格,并合并相邻的同一标签。""" - runs: list[dict[str, Any]] = [] - for index, cell in enumerate(cells): - if runs and runs[-1]["label"] == cell["label"]: - runs[-1]["end"] = cell["end"] - runs[-1]["indices"].append(index) - else: - runs.append( - { - "start": cell["start"], - "end": cell["end"], - "label": cell["label"], - "indices": [index], - } + def _map_cluster_centers( + self, + session_id: str, + cluster_centers: Any, + ) -> list[dict[str, Any]]: + """Map FunASR's temporary cluster labels onto stable, capped session IDs.""" + import numpy as np + + clusters = self.speaker_clusters.setdefault(session_id, []) + used_ids: set[int] = set() + mapped: list[dict[str, Any]] = [] + + for raw_center in cluster_centers: + center = self._normalize_embedding(raw_center) + available = [ + cluster for cluster in clusters + if int(cluster["speaker_id"]) not in used_ids + ] + best_cluster = max( + available, + key=lambda cluster: float(np.dot(center, cluster["embedding"])), + default=None, + ) + best_score = ( + float(np.dot(center, best_cluster["embedding"])) + if best_cluster is not None + else -1.0 + ) + matched = ( + best_cluster is not None + and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD + ) + created = False + + if not matched and len(clusters) < ONLINE_MAX_SPEAKERS: + speaker_id = len(clusters) + clusters.append( + {"speaker_id": speaker_id, "embedding": center, "count": 1} ) - return runs + confidence = 0.75 + strategy = "online_embedding_cluster_new" + created = True + else: + if best_cluster is None: + # If this turn has more temporary clusters than available IDs, + # follow FunASR and fall back to the nearest existing identity. + best_cluster = max( + clusters, + key=lambda cluster: float(np.dot(center, cluster["embedding"])), + ) + best_score = float(np.dot(center, best_cluster["embedding"])) + speaker_id = int(best_cluster["speaker_id"]) + confidence = max(0.0, best_score) + if matched: + count = int(best_cluster["count"]) + weight = 1.0 / min(count + 1, 20) + best_cluster["embedding"] = self._normalize_embedding( + best_cluster["embedding"] * (1.0 - weight) + + center * weight + ) + best_cluster["count"] = count + 1 + strategy = "online_embedding_cluster_match" + else: + # Once the configured identity limit is reached, keep IDs stable + # by assigning unmatched clusters to their nearest known center. + strategy = "online_embedding_cluster_limit_fallback" + + used_ids.add(speaker_id) + mapped.append( + { + "speaker_id": speaker_id, + "speaker_confidence": round( + max(0.6, min(1.0, confidence)), 3 + ), + "speaker_strategy": strategy, + } + ) + return mapped + + def _cluster_speaker_history( + self, + session_id: str, + ) -> tuple[list[list[float]], list[dict[str, Any]]]: + """Re-cluster the rolling CAM++ history with FunASR's own backend.""" + import numpy as np + import torch + from funasr.models.campplus.cluster_backend import ClusterBackend + from funasr.models.campplus.utils import postprocess + + history = self.speaker_history[session_id] + embeddings = torch.as_tensor( + np.stack(list(history["embeddings"])), + dtype=torch.float32, + device="cpu", + ) + if self._speaker_cluster_backend is None: + self._speaker_cluster_backend = ClusterBackend( + merge_thr=ONLINE_CLUSTER_MERGE_THRESHOLD + ).to("cpu") + self._speaker_postprocess = postprocess + + # ClusterBackend produces turn-local labels; postprocess also aligns overlap + # boundaries and smooths short speaker runs before stable IDs are assigned. + labels = self._speaker_cluster_backend(embeddings, oracle_num=None) + labels = np.asarray(labels) + chunks = [ + [start_ms / 1000.0, end_ms / 1000.0, None] + for start_ms, end_ms in history["chunks"] + ] + segments, centers = self._speaker_postprocess( + chunks, + None, + labels, + embeddings, + return_spk_center=True, + ) + stable_clusters = self._map_cluster_centers(session_id, centers) + return segments, stable_clusters def _assign_embedding( self, @@ -529,7 +638,7 @@ class AuxiliaryRuntime: start_time_ms: float, end_time_ms: float, ) -> dict[str, Any]: - """将一个稳定说话人段映射到会话中心,并只在段完成后更新中心。""" + """Assign a fallback whole-turn embedding with FunASR's 15-speaker policy.""" import numpy as np embedding = self._normalize_embedding(embedding) @@ -546,19 +655,25 @@ class AuxiliaryRuntime: if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: count = int(best_cluster["count"]) + weight = 1.0 / min(count + 1, 20) best_cluster["embedding"] = self._normalize_embedding( - (best_cluster["embedding"] * count) + embedding + best_cluster["embedding"] * (1.0 - weight) + + embedding * weight ) best_cluster["count"] = count + 1 speaker_id = int(best_cluster["speaker_id"]) confidence = best_score strategy = "online_embedding_cluster_match" - else: - # 不设置会话人数上限;与已有中心不匹配的声纹建立新身份。 + elif len(clusters) < ONLINE_MAX_SPEAKERS: speaker_id = len(clusters) clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1}) confidence = 0.75 strategy = "online_embedding_cluster_new" + else: + # Match FunASR's fallback after the identity limit is reached. + speaker_id = int(best_cluster["speaker_id"]) if best_cluster else 0 + confidence = max(0.0, best_score) + strategy = "online_embedding_cluster_limit_fallback" return { "speaker_id": speaker_id, @@ -567,7 +682,7 @@ class AuxiliaryRuntime: "speaker_confidence": round(max(0.6, min(1.0, confidence)), 3), "speaker_strategy": strategy, "speaker_status": "confirmed", - "speaker_reason": "CAM++ 滑窗声纹已完成会话内在线聚类", + "speaker_reason": "CAM++ rolling history was clustered with FunASR's speaker backend", "start_time": start_time_ms, "end_time": end_time_ms, } @@ -579,9 +694,7 @@ class AuxiliaryRuntime: start_time_ms: float, end_time_ms: float, ) -> list[dict[str, Any]]: - """用重叠 CAM++ 窗口跟踪一个 FunASR VAD turn 内的说话人变化。""" - import numpy as np - + """Track a completed VAD turn using FunASR's rolling-window clusterer.""" async with self.inference_lock: now = time.monotonic() for stale_id, seen in list(self.speaker_last_seen.items()): @@ -592,114 +705,46 @@ class AuxiliaryRuntime: if session_id not in self.speaker_last_seen or not windows: return [] - # 每个 turn 先在局部聚类;只对平滑后留下的长段更新持久中心, - # 避免一个短暂的误判污染后续 turn 的说话人身份。 - global_clusters = self.speaker_clusters.get(session_id, []) - local_centers: list[dict[str, Any]] = [] - cells: list[dict[str, Any]] = [] - for index, (window_start, window_end, embedding) in enumerate(windows): - best_label: tuple[str, int] | None = None - best_score = -1.0 - for cluster in global_clusters: - score = float(np.dot(embedding, cluster["embedding"])) - if score > best_score: - best_score = score - best_label = ("global", int(cluster["speaker_id"])) - if best_label is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: - label = best_label - else: - local_index = None - local_score = -1.0 - for candidate_index, candidate in enumerate(local_centers): - score = float(np.dot(embedding, candidate["embedding"])) - if score > local_score: - local_score = score - local_index = candidate_index - if local_index is not None and local_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: - candidate = local_centers[local_index] - count = int(candidate["count"]) - candidate["embedding"] = self._normalize_embedding( - candidate["embedding"] * count + embedding - ) - candidate["count"] = count + 1 - label = ("local", local_index) - else: - local_index = len(local_centers) - local_centers.append( - {"embedding": embedding.copy(), "count": 1} - ) - label = ("local", local_index) - cells.append( + history = self.speaker_history.get(session_id) + if history is None: + history = { + "chunks": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS), + "embeddings": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS), + } + self.speaker_history[session_id] = history + + # Keep absolute timestamps so each turn can be re-clustered against recent turns. + for window_start, window_end, embedding in windows: + history["chunks"].append( + (start_time_ms + window_start, start_time_ms + window_end) + ) + history["embeddings"].append(embedding.copy()) + + clustered_segments, stable_clusters = self._cluster_speaker_history(session_id) + results: list[dict[str, Any]] = [] + for segment_start, segment_end, cluster_id in clustered_segments: + segment_start_ms = max(start_time_ms, float(segment_start) * 1000.0) + segment_end_ms = min(end_time_ms, float(segment_end) * 1000.0) + if segment_end_ms <= segment_start_ms: + continue + cluster_index = int(cluster_id) + if cluster_index < 0 or cluster_index >= len(stable_clusters): + continue + stable = stable_clusters[cluster_index] + speaker_id = int(stable["speaker_id"]) + results.append( { - "start": float(window_start), - "end": float(window_end), - "center": (window_start + window_end) / 2, - "label": label, - "window_index": index, + "speaker_id": speaker_id, + "speaker_name": f"说话人 {speaker_id + 1}", + "speaker_evidence": "fresh", + "speaker_confidence": stable["speaker_confidence"], + "speaker_strategy": stable["speaker_strategy"], + "speaker_status": "confirmed", + "speaker_reason": "CAM++ rolling history was clustered with FunASR's speaker backend", + "start_time": segment_start_ms, + "end_time": segment_end_ms, } ) - - # 用相邻窗口中心的中点确定切换边界,延续 FunASR 滑窗跟踪的时间语义。 - turn_duration_ms = max(0.0, end_time_ms - start_time_ms, windows[-1][1]) - for index, cell in enumerate(cells): - cell["start"] = ( - 0.0 - if index == 0 - else (cells[index - 1]["center"] + cell["center"]) / 2 - ) - cell["end"] = ( - turn_duration_ms - if index == len(cells) - 1 - else (cell["center"] + cells[index + 1]["center"]) / 2 - ) - - min_segment_ms = max( - 1500, int(os.getenv("FUNASR_SPEAKER_MIN_SEGMENT_MS", "3000")) - ) - while len(cells) > 1: - runs = self._speaker_runs(cells) - short_run = next( - ( - index - for index, run in enumerate(runs) - if run["end"] - run["start"] < min_segment_ms - ), - None, - ) - if short_run is None: - break - run = runs[short_run] - if short_run == 0: - target_label = runs[1]["label"] - elif short_run == len(runs) - 1: - target_label = runs[-2]["label"] - else: - previous = runs[short_run - 1] - following = runs[short_run + 1] - target_label = ( - previous["label"] - if previous["end"] - previous["start"] - >= following["end"] - following["start"] - else following["label"] - ) - for cell in cells: - if run["start"] <= cell["start"] < run["end"]: - cell["label"] = target_label - - results: list[dict[str, Any]] = [] - for run in self._speaker_runs(cells): - vectors = [ - windows[cells[index]["window_index"]][2] - for index in run["indices"] - ] - mean_embedding = self._normalize_embedding(np.mean(vectors, axis=0)) - result = self._assign_embedding( - session_id, - mean_embedding, - start_time_ms + run["start"], - min(end_time_ms, start_time_ms + run["end"]), - ) - results.append(result) return results async def resolve_speaker( @@ -731,6 +776,7 @@ class AuxiliaryRuntime: def reset_speaker_session(self, session_id: str) -> None: """释放已结束 WebSocket 的聚类中心,防止长时间运行时内存增长。""" self.speaker_clusters.pop(session_id, None) + self.speaker_history.pop(session_id, None) self.speaker_last_seen.pop(session_id, None)