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 FUNASR_SPEAKER_HOP_MS=750
# CAM++ 每批最多计算多少个窗口,避免很长 turn 一次性占用过多显存。 # CAM++ 每批最多计算多少个窗口,避免很长 turn 一次性占用过多显存。
FUNASR_SPEAKER_BATCH_SIZE=16 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 FUNASR_SPEAKER_MIN_SEGMENT_MS=3000

View File

@ -9,6 +9,7 @@ import os
import sys import sys
import tempfile import tempfile
import time import time
from collections import deque
from collections.abc import Mapping from collections.abc import Mapping
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@ -30,7 +31,11 @@ AUXILIARY_PORT = 8010
load_dotenv(PROJECT_ROOT / ".env") load_dotenv(PROJECT_ROOT / ".env")
MODELS_DIR = Path(os.getenv("MODEL_DIR", str(PROJECT_ROOT / "models"))).resolve() MODELS_DIR = Path(os.getenv("MODEL_DIR", str(PROJECT_ROOT / "models"))).resolve()
AUXILIARY_DEVICE = os.getenv("AUXILIARY_DEVICE", "cuda:0") 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 MIN_ONLINE_SPEAKER_AUDIO_MS = 800
# Realtime VAD is loaded by the WebSocket process. This service loads CAM++. # Realtime VAD is loaded by the WebSocket process. This service loads CAM++.
# Its standalone VAD HTTP endpoint remains available on demand. # Its standalone VAD HTTP endpoint remains available on demand.
@ -94,7 +99,11 @@ class AuxiliaryRuntime:
self.inference_lock = asyncio.Lock() self.inference_lock = asyncio.Lock()
# 每个 WebSocket session 独立维护聚类中心,避免不同浏览器会话互相污染。 # 每个 WebSocket session 独立维护聚类中心,避免不同浏览器会话互相污染。
self.speaker_clusters: dict[str, list[dict[str, Any]]] = {} 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_last_seen: dict[str, float] = {}
self._speaker_cluster_backend: Any | None = None
self._speaker_postprocess: Any | None = None
def _preload_kinds(self) -> set[str]: def _preload_kinds(self) -> set[str]:
"""读取需要在启动时加载的模型类型,默认不加载完整 diarization。""" """读取需要在启动时加载的模型类型,默认不加载完整 diarization。"""
@ -503,24 +512,124 @@ class AuxiliaryRuntime:
for (start_ms, end_ms), embedding in zip(spans, embeddings) for (start_ms, end_ms), embedding in zip(spans, embeddings)
] ]
@staticmethod def _map_cluster_centers(
def _speaker_runs(cells: list[dict[str, Any]]) -> list[dict[str, Any]]: self,
"""把重叠声纹窗变成连续时间格,并合并相邻的同一标签。""" session_id: str,
runs: list[dict[str, Any]] = [] cluster_centers: Any,
for index, cell in enumerate(cells): ) -> list[dict[str, Any]]:
if runs and runs[-1]["label"] == cell["label"]: """Map FunASR's temporary cluster labels onto stable, capped session IDs."""
runs[-1]["end"] = cell["end"] import numpy as np
runs[-1]["indices"].append(index)
else: clusters = self.speaker_clusters.setdefault(session_id, [])
runs.append( used_ids: set[int] = set()
{ mapped: list[dict[str, Any]] = []
"start": cell["start"],
"end": cell["end"], for raw_center in cluster_centers:
"label": cell["label"], center = self._normalize_embedding(raw_center)
"indices": [index], 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( def _assign_embedding(
self, self,
@ -529,7 +638,7 @@ class AuxiliaryRuntime:
start_time_ms: float, start_time_ms: float,
end_time_ms: float, end_time_ms: float,
) -> dict[str, Any]: ) -> dict[str, Any]:
"""将一个稳定说话人段映射到会话中心,并只在段完成后更新中心。""" """Assign a fallback whole-turn embedding with FunASR's 15-speaker policy."""
import numpy as np import numpy as np
embedding = self._normalize_embedding(embedding) embedding = self._normalize_embedding(embedding)
@ -546,19 +655,25 @@ class AuxiliaryRuntime:
if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD:
count = int(best_cluster["count"]) count = int(best_cluster["count"])
weight = 1.0 / min(count + 1, 20)
best_cluster["embedding"] = self._normalize_embedding( best_cluster["embedding"] = self._normalize_embedding(
(best_cluster["embedding"] * count) + embedding best_cluster["embedding"] * (1.0 - weight)
+ embedding * weight
) )
best_cluster["count"] = count + 1 best_cluster["count"] = count + 1
speaker_id = int(best_cluster["speaker_id"]) speaker_id = int(best_cluster["speaker_id"])
confidence = best_score confidence = best_score
strategy = "online_embedding_cluster_match" strategy = "online_embedding_cluster_match"
else: elif len(clusters) < ONLINE_MAX_SPEAKERS:
# 不设置会话人数上限;与已有中心不匹配的声纹建立新身份。
speaker_id = len(clusters) speaker_id = len(clusters)
clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1}) clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1})
confidence = 0.75 confidence = 0.75
strategy = "online_embedding_cluster_new" 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 { return {
"speaker_id": speaker_id, "speaker_id": speaker_id,
@ -567,7 +682,7 @@ class AuxiliaryRuntime:
"speaker_confidence": round(max(0.6, min(1.0, confidence)), 3), "speaker_confidence": round(max(0.6, min(1.0, confidence)), 3),
"speaker_strategy": strategy, "speaker_strategy": strategy,
"speaker_status": "confirmed", "speaker_status": "confirmed",
"speaker_reason": "CAM++ 滑窗声纹已完成会话内在线聚类", "speaker_reason": "CAM++ rolling history was clustered with FunASR's speaker backend",
"start_time": start_time_ms, "start_time": start_time_ms,
"end_time": end_time_ms, "end_time": end_time_ms,
} }
@ -579,9 +694,7 @@ class AuxiliaryRuntime:
start_time_ms: float, start_time_ms: float,
end_time_ms: float, end_time_ms: float,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""用重叠 CAM++ 窗口跟踪一个 FunASR VAD turn 内的说话人变化。""" """Track a completed VAD turn using FunASR's rolling-window clusterer."""
import numpy as np
async with self.inference_lock: async with self.inference_lock:
now = time.monotonic() now = time.monotonic()
for stale_id, seen in list(self.speaker_last_seen.items()): 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: if session_id not in self.speaker_last_seen or not windows:
return [] return []
# 每个 turn 先在局部聚类;只对平滑后留下的长段更新持久中心, history = self.speaker_history.get(session_id)
# 避免一个短暂的误判污染后续 turn 的说话人身份。 if history is None:
global_clusters = self.speaker_clusters.get(session_id, []) history = {
local_centers: list[dict[str, Any]] = [] "chunks": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS),
cells: list[dict[str, Any]] = [] "embeddings": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS),
for index, (window_start, window_end, embedding) in enumerate(windows): }
best_label: tuple[str, int] | None = None self.speaker_history[session_id] = history
best_score = -1.0
for cluster in global_clusters: # Keep absolute timestamps so each turn can be re-clustered against recent turns.
score = float(np.dot(embedding, cluster["embedding"])) for window_start, window_end, embedding in windows:
if score > best_score: history["chunks"].append(
best_score = score (start_time_ms + window_start, start_time_ms + window_end)
best_label = ("global", int(cluster["speaker_id"])) )
if best_label is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: history["embeddings"].append(embedding.copy())
label = best_label
else: clustered_segments, stable_clusters = self._cluster_speaker_history(session_id)
local_index = None results: list[dict[str, Any]] = []
local_score = -1.0 for segment_start, segment_end, cluster_id in clustered_segments:
for candidate_index, candidate in enumerate(local_centers): segment_start_ms = max(start_time_ms, float(segment_start) * 1000.0)
score = float(np.dot(embedding, candidate["embedding"])) segment_end_ms = min(end_time_ms, float(segment_end) * 1000.0)
if score > local_score: if segment_end_ms <= segment_start_ms:
local_score = score continue
local_index = candidate_index cluster_index = int(cluster_id)
if local_index is not None and local_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: if cluster_index < 0 or cluster_index >= len(stable_clusters):
candidate = local_centers[local_index] continue
count = int(candidate["count"]) stable = stable_clusters[cluster_index]
candidate["embedding"] = self._normalize_embedding( speaker_id = int(stable["speaker_id"])
candidate["embedding"] * count + embedding results.append(
)
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(
{ {
"start": float(window_start), "speaker_id": speaker_id,
"end": float(window_end), "speaker_name": f"说话人 {speaker_id + 1}",
"center": (window_start + window_end) / 2, "speaker_evidence": "fresh",
"label": label, "speaker_confidence": stable["speaker_confidence"],
"window_index": index, "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 return results
async def resolve_speaker( async def resolve_speaker(
@ -731,6 +776,7 @@ class AuxiliaryRuntime:
def reset_speaker_session(self, session_id: str) -> None: def reset_speaker_session(self, session_id: str) -> None:
"""释放已结束 WebSocket 的聚类中心,防止长时间运行时内存增长。""" """释放已结束 WebSocket 的聚类中心,防止长时间运行时内存增长。"""
self.speaker_clusters.pop(session_id, None) self.speaker_clusters.pop(session_id, None)
self.speaker_history.pop(session_id, None)
self.speaker_last_seen.pop(session_id, None) self.speaker_last_seen.pop(session_id, None)