Align speaker tracking with FunASR clustering

main
Bifang 2026-09-24 10:35:47 +08:00
parent 14420427a9
commit 5c58759099
2 changed files with 183 additions and 131 deletions

View File

@ -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

View File

@ -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)
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}
)
confidence = 0.75
strategy = "online_embedding_cluster_new"
created = True
else:
runs.append(
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(
{
"start": cell["start"],
"end": cell["end"],
"label": cell["label"],
"indices": [index],
"speaker_id": speaker_id,
"speaker_confidence": round(
max(0.6, min(1.0, confidence)), 3
),
"speaker_strategy": strategy,
}
)
return runs
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
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)
)
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["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)