Align speaker tracking with FunASR clustering
parent
14420427a9
commit
5c58759099
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue