Improve FunASR realtime speaker tracking

main
Bifang 2026-09-24 09:35:55 +08:00
parent 1a239cffbf
commit 14420427a9
5 changed files with 622 additions and 78 deletions

View File

@ -13,6 +13,13 @@ FUNASR_ENCODER_LOOK_BACK=4
FUNASR_DECODER_LOOK_BACK=1 FUNASR_DECODER_LOOK_BACK=1
FUNASR_VAD_CHUNK_MS=200 FUNASR_VAD_CHUNK_MS=200
FUNASR_MAX_SEGMENT_SEC=30 FUNASR_MAX_SEGMENT_SEC=30
# CAM++ 段内说话人追踪沿用 FunASR 的重叠短窗方案:1.5 秒窗、0.75 秒步长。
FUNASR_SPEAKER_WINDOW_MS=1500
FUNASR_SPEAKER_HOP_MS=750
# CAM++ 每批最多计算多少个窗口,避免很长 turn 一次性占用过多显存。
FUNASR_SPEAKER_BATCH_SIZE=16
# 小于该时长的相邻说话人段会合并,避免短窗噪声把一句话切得过碎。
FUNASR_SPEAKER_MIN_SEGMENT_MS=3000
# CAM++ is required and started by the backend launcher. # CAM++ is required and started by the backend launcher.
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010 AUXILIARY_SERVICE_URL=http://127.0.0.1:8010

View File

@ -376,30 +376,162 @@ class AuxiliaryRuntime:
output = self._run_embedding_pipeline(model_pipeline, audio_path) output = self._run_embedding_pipeline(model_pipeline, audio_path)
return self._normalize_embedding(output) return self._normalize_embedding(output)
async def resolve_speaker( @staticmethod
self, def _embedding_rows(result: Any, expected_count: int) -> Any:
audio_path: str, """兼容 ModelScope 批量 embedding 的 Tensor、数组和逐条结果格式。"""
session_id: str, import numpy as np
start_time_ms: float,
end_time_ms: float, values = AuxiliaryRuntime._extract_embedding_value(result)
) -> dict[str, Any] | None: if values is None:
"""对一个实时 turn 提取声纹,并更新该 session 的在线聚类中心。""" raise RuntimeError("speaker pipeline returned no batch embeddings")
async with self.inference_lock: detach = getattr(values, "detach", None)
# 异常断网时客户端可能来不及 reset,过期状态在下一次请求时回收。 if callable(detach):
now = time.monotonic() values = detach()
for stale_id, seen in list(self.speaker_last_seen.items()): cpu = getattr(values, "cpu", None)
if now - seen > 1800: if callable(cpu):
self.reset_speaker_session(stale_id) values = cpu()
self.speaker_last_seen[session_id] = now numpy_method = getattr(values, "numpy", None)
embedding = await asyncio.to_thread(self._extract_embedding_sync, audio_path) if callable(numpy_method):
if embedding is None: values = numpy_method()
return {"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0, try:
"speaker_status": "insufficient_audio", "speaker_reason": "音频不足 800ms,未提取声纹"} array = np.asarray(values, dtype=np.float32)
except (TypeError, ValueError):
array = np.empty((0, 0), dtype=np.float32)
if array.ndim == 3 and array.shape[0] == expected_count and array.shape[1] == 1:
array = array[:, 0, :]
if array.ndim == 1 and expected_count == 1:
array = array.reshape(1, -1)
if array.ndim == 2 and array.shape[0] == expected_count:
rows = [AuxiliaryRuntime._normalize_embedding(row) for row in array]
return rows
# 一些 ModelScope 版本返回每条音频一个对象,而不是一张二维矩阵。
if isinstance(values, (list, tuple)) and len(values) == expected_count:
rows = []
for item in values:
value = AuxiliaryRuntime._extract_embedding_value(item)
if value is None:
raise RuntimeError("speaker pipeline returned an item without an embedding")
rows.append(AuxiliaryRuntime._normalize_embedding(value))
return rows
raise RuntimeError(
"speaker pipeline batch size mismatch: "
f"expected {expected_count}, received shape {getattr(array, 'shape', None)}"
)
@staticmethod
def _write_pcm_window(path: str, samples: Any) -> None:
"""把补齐后的 16kHz mono 浮点窗写成 CAM++ pipeline 可读的 PCM WAV。"""
import numpy as np
import wave
pcm = (np.clip(samples, -1.0, 1.0) * 32767.0).astype("<i2")
with wave.open(path, "wb") as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(16000)
wav_file.writeframes(pcm.tobytes())
def _extract_window_embeddings_sync(self, audio_path: str) -> list[tuple[float, float, Any]]:
"""按 FunASR CAM++ 的 1.5s/0.75s 重叠窗批量提取声纹。"""
import librosa
import numpy as np
audio, _ = librosa.load(audio_path, sr=16000, mono=True)
audio = np.asarray(audio, dtype=np.float32).reshape(-1)
min_samples = int(16000 * MIN_ONLINE_SPEAKER_AUDIO_MS / 1000)
if audio.size < min_samples:
return []
window_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_WINDOW_MS", "1500")))
hop_ms = max(250, int(os.getenv("FUNASR_SPEAKER_HOP_MS", "750")))
window_samples = int(16000 * window_ms / 1000)
hop_samples = int(16000 * hop_ms / 1000)
model_id = self._speaker_embedding_model_id()
if model_id is None:
raise RuntimeError("no loaded CAM++ speaker verification model is available")
model_pipeline = self.models[model_id]
spans: list[tuple[float, float]] = []
paths: list[str] = []
with tempfile.TemporaryDirectory(prefix="funasr-campp-windows-") as temp_dir:
last_end = 0
for suggested_start in range(0, audio.size, hop_samples):
end = min(suggested_start + window_samples, audio.size)
if end <= last_end:
break
last_end = end
# 和 FunASR sv_chunk 一样把尾窗右对齐,确保不足 1.5 秒的末尾窗
# 尽量包含完整的新语音,而不是只用较短尾音再补大量零。
start = max(0, end - window_samples)
chunk = audio[start:end]
if chunk.size < min_samples:
break
# 忽略近静音窗,避免把背景底噪添加成一个新说话人。
rms = float(np.sqrt(np.mean(np.square(chunk)))) if chunk.size else 0.0
if rms >= 0.0015:
padded = np.zeros(window_samples, dtype=np.float32)
padded[: chunk.size] = chunk
path = str(Path(temp_dir) / f"window-{len(paths):04d}.wav")
self._write_pcm_window(path, padded)
paths.append(path)
spans.append((start * 1000 / 16000, end * 1000 / 16000))
if end >= audio.size:
break
if not paths:
return []
embeddings = []
batch_size = max(1, min(64, int(os.getenv("FUNASR_SPEAKER_BATCH_SIZE", "16"))))
for offset in range(0, len(paths), batch_size):
path_batch = paths[offset : offset + batch_size]
try:
batch_result = model_pipeline(path_batch, output_emb=True)
batch_embeddings = self._embedding_rows(batch_result, len(path_batch))
except Exception:
# 有些旧 pipeline 不接受多文件批次;逐窗调用保持同一预处理路径。
batch_embeddings = [
self._normalize_embedding(
self._run_embedding_pipeline(model_pipeline, path)
)
for path in path_batch
]
embeddings.extend(batch_embeddings)
return [
(start_ms, end_ms, embedding)
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],
}
)
return runs
def _assign_embedding(
self,
session_id: str,
embedding: Any,
start_time_ms: float,
end_time_ms: float,
) -> dict[str, Any]:
"""将一个稳定说话人段映射到会话中心,并只在段完成后更新中心。"""
import numpy as np import numpy as np
# reset 可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。
if session_id not in self.speaker_last_seen:
return None
embedding = self._normalize_embedding(embedding) embedding = self._normalize_embedding(embedding)
clusters = self.speaker_clusters.setdefault(session_id, []) clusters = self.speaker_clusters.setdefault(session_id, [])
best_cluster: dict[str, Any] | None = None best_cluster: dict[str, Any] | None = None
@ -422,6 +554,7 @@ class AuxiliaryRuntime:
confidence = best_score confidence = best_score
strategy = "online_embedding_cluster_match" strategy = "online_embedding_cluster_match"
else: else:
# 不设置会话人数上限;与已有中心不匹配的声纹建立新身份。
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
@ -434,11 +567,167 @@ 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": "当前片段独立声纹已完成在线聚类", "speaker_reason": "CAM++ 滑窗声纹已完成会话内在线聚类",
"start_time": start_time_ms, "start_time": start_time_ms,
"end_time": end_time_ms, "end_time": end_time_ms,
} }
async def track_speakers(
self,
audio_path: str,
session_id: str,
start_time_ms: float,
end_time_ms: float,
) -> list[dict[str, Any]]:
"""用重叠 CAM++ 窗口跟踪一个 FunASR VAD turn 内的说话人变化。"""
import numpy as np
async with self.inference_lock:
now = time.monotonic()
for stale_id, seen in list(self.speaker_last_seen.items()):
if now - seen > 1800:
self.reset_speaker_session(stale_id)
self.speaker_last_seen[session_id] = now
windows = await asyncio.to_thread(self._extract_window_embeddings_sync, audio_path)
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(
{
"start": float(window_start),
"end": float(window_end),
"center": (window_start + window_end) / 2,
"label": label,
"window_index": index,
}
)
# 用相邻窗口中心的中点确定切换边界,延续 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(
self,
audio_path: str,
session_id: str,
start_time_ms: float,
end_time_ms: float,
) -> dict[str, Any] | None:
"""对一个实时 turn 提取声纹,并更新该 session 的在线聚类中心。"""
async with self.inference_lock:
# 异常断网时客户端可能来不及 reset,过期状态在下一次请求时回收。
now = time.monotonic()
for stale_id, seen in list(self.speaker_last_seen.items()):
if now - seen > 1800:
self.reset_speaker_session(stale_id)
self.speaker_last_seen[session_id] = now
embedding = await asyncio.to_thread(self._extract_embedding_sync, audio_path)
if embedding is None:
return {"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0,
"speaker_status": "insufficient_audio", "speaker_reason": "音频不足 800ms,未提取声纹"}
# reset 可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。
if session_id not in self.speaker_last_seen:
return None
return self._assign_embedding(
session_id, embedding, start_time_ms, end_time_ms
)
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)
@ -627,6 +916,43 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
Path(temp_path).unlink(missing_ok=True) Path(temp_path).unlink(missing_ok=True)
async def speaker_track_handler(request: web.Request) -> web.Response:
"""接收一个 FunASR turn,返回经过短段平滑的 CAM++ 滑窗标签。"""
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
form = await request.post()
upload = form.get("file")
if not isinstance(upload, FileField):
return web.json_response({"error": "multipart field 'file' is required"}, status=400)
session_id = str(form.get("session_id") or "").strip()
if not session_id:
return web.json_response({"error": "multipart field 'session_id' is required"}, status=400)
try:
start_time_ms = _parse_form_float(form.get("start_time_ms"), "start_time_ms", default=0.0)
end_time_ms = _parse_form_float(form.get("end_time_ms"), "end_time_ms", default=start_time_ms)
except ValueError:
return web.json_response({"error": "turn time fields must be numbers"}, status=400)
temp_path: str | None = None
try:
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as temp_file:
temp_file.write(upload.file.read())
temp_path = temp_file.name
segments = await runtime.track_speakers(
temp_path, session_id, start_time_ms, end_time_ms
)
return web.json_response({"segments": segments})
except Exception as exc:
print(
f"[speaker] track failed: session_id={session_id}, "
f"start={start_time_ms}, end={end_time_ms}, error={exc}",
flush=True,
)
return web.json_response({"error": str(exc)}, status=500)
finally:
if temp_path:
Path(temp_path).unlink(missing_ok=True)
async def speaker_reset_handler(request: web.Request) -> web.Response: async def speaker_reset_handler(request: web.Request) -> web.Response:
"""释放已经结束的实时会话聚类中心。""" """释放已经结束的实时会话聚类中心。"""
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY] runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
@ -648,6 +974,7 @@ async def create_app() -> web.Application:
app.router.add_post("/v1/punctuation", punctuation_handler) app.router.add_post("/v1/punctuation", punctuation_handler)
app.router.add_post("/v1/diarization", diarization_handler) app.router.add_post("/v1/diarization", diarization_handler)
app.router.add_post("/v1/speaker/resolve", speaker_resolve_handler) app.router.add_post("/v1/speaker/resolve", speaker_resolve_handler)
app.router.add_post("/v1/speaker/track", speaker_track_handler)
app.router.add_post("/v1/speaker/reset", speaker_reset_handler) app.router.add_post("/v1/speaker/reset", speaker_reset_handler)
return app return app

View File

@ -116,6 +116,35 @@ class AuxiliaryModelService:
# 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。 # 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。
return decoded return decoded
async def track_speakers(
self,
pcm_bytes: bytes,
session_id: str,
start_time_ms: float,
end_time_ms: float,
) -> list[dict[str, Any]]:
"""对已完成的 FunASR turn 按 CAM++ 滑窗追踪段内说话人。"""
if self._session is None:
raise RuntimeError("auxiliary model service is not started")
form = FormData()
form.add_field("file", pcm16_to_wav(pcm_bytes), filename="turn.wav", content_type="audio/wav")
form.add_field("session_id", session_id)
form.add_field("start_time_ms", str(start_time_ms))
form.add_field("end_time_ms", str(end_time_ms))
endpoint = self.config.base_url.rstrip("/") + "/v1/speaker/track"
async with self._session.post(endpoint, data=form) as response:
body = await response.text()
if response.status >= 400:
raise RuntimeError(f"auxiliary speaker tracking failed ({response.status}): {body[:500]}")
try:
decoded = await response.json(content_type=None)
except ValueError as exc:
raise RuntimeError(f"auxiliary speaker tracking returned invalid JSON: {body[:500]}") from exc
segments = decoded.get("segments", []) if isinstance(decoded, dict) else []
if not isinstance(segments, list):
return []
return [segment for segment in segments if isinstance(segment, dict)]
async def reset_speaker_session(self, session_id: str) -> None: async def reset_speaker_session(self, session_id: str) -> None:
"""通知辅助服务释放当前 WebSocket 对应的在线聚类状态。""" """通知辅助服务释放当前 WebSocket 对应的在线聚类状态。"""
if self._session is None: if self._session is None:

View File

@ -650,6 +650,20 @@ async def ws_serve(websocket, path=None):
record_error(f"vad inference failed: {e}") record_error(f"vad inference failed: {e}")
speech_start_i, speech_end_i = -1, -1 speech_start_i, speech_end_i = -1, -1
# 把 FunASR VAD 的绝对音频时间发给桥接层,用于去掉首尾静音,
# 让 CAM++ 滑窗只分析当前 VAD turn,而不是整段会话的缓冲音频。
if speech_start_i != -1 or speech_end_i != -1:
await websocket.send(
json.dumps(
{
"event": "vad",
"speech_start_ms": speech_start_i if speech_start_i != -1 else None,
"speech_end_ms": speech_end_i if speech_end_i != -1 else None,
},
ensure_ascii=False,
)
)
if speech_start_i != -1: if speech_start_i != -1:
speech_start = True speech_start = True
if duration_ms > 0: if duration_ms > 0:
@ -890,15 +904,18 @@ async def async_asr_online(websocket, audio_in: bytes):
if websocket.mode == "2pass" and websocket.status_dict_asr_online.get("is_final", False): if websocket.mode == "2pass" and websocket.status_dict_asr_online.get("is_final", False):
return return
if rec_result.get("text"): is_final = bool(
websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking)
)
# 即使最终解码没有新增字符,也必须显式发送 final 事件;桥接层靠它
# 结束静音期间的 turn,否则最后一条文本会一直停留在 interim 状态。
if rec_result.get("text") or is_final:
mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode
message = { message = {
"mode": mode, "mode": mode,
"text": rec_result["text"], "text": rec_result.get("text", ""),
"wav_name": websocket.wav_name, "wav_name": websocket.wav_name,
"is_final": bool( "is_final": is_final,
websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking)
),
} }
await websocket.send(json.dumps(message, ensure_ascii=False)) await websocket.send(json.dumps(message, ensure_ascii=False))

View File

@ -44,6 +44,7 @@ FRAME_BYTES = max(2, round(60 * CHUNK_SIZE[1] / CHUNK_INTERVAL * PCM_BYTES_PER_M
MAX_SPEAKER_AUDIO_BYTES = 60 * SAMPLE_RATE * 2 MAX_SPEAKER_AUDIO_BYTES = 60 * SAMPLE_RATE * 2
MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2) MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2)
FINALIZE_TIMEOUT_SECONDS = max(30, int(os.getenv("FUNASR_FINALIZE_TIMEOUT_SECONDS", "300"))) FINALIZE_TIMEOUT_SECONDS = max(30, int(os.getenv("FUNASR_FINALIZE_TIMEOUT_SECONDS", "300")))
TURN_SENTENCE_ID_STRIDE = 100
AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService) AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
SESSION_REGISTRY_KEY = web.AppKey("sessions", dict) SESSION_REGISTRY_KEY = web.AppKey("sessions", dict)
@ -121,6 +122,96 @@ class SpeakerJob:
end_time_ms: float end_time_ms: float
def split_text_by_speaker_segments(
text: str,
segments: list[dict[str, Any]],
turn_start_ms: float,
turn_end_ms: float,
) -> list[dict[str, Any]]:
"""按 CAM++ 时间段近似拆分 ASR 文本,并保留每段稳定说话人标签。"""
usable: list[dict[str, Any]] = []
for segment in sorted(
segments,
key=lambda item: float(item.get("start_time", turn_start_ms)),
):
try:
start = max(turn_start_ms, float(segment.get("start_time", turn_start_ms)))
end = min(turn_end_ms, float(segment.get("end_time", turn_end_ms)))
speaker_id = int(segment.get("speaker_id", -1))
except (TypeError, ValueError):
continue
if end <= start or speaker_id < 0:
continue
speaker = {
"speaker_id": speaker_id,
"speaker_name": str(segment.get("speaker_name") or f"说话人 {speaker_id + 1}"),
"speaker_confidence": float(segment.get("speaker_confidence") or 0),
"speaker_status": str(segment.get("speaker_status") or "confirmed"),
}
if usable and usable[-1]["speaker"]["speaker_id"] == speaker_id:
usable[-1]["end_time_ms"] = end
else:
usable.append(
{"start_time_ms": start, "end_time_ms": end, "speaker": speaker}
)
if not usable:
return []
if len(usable) == 1:
usable[0]["text"] = text
return usable
# 双重保护:服务端有最短段长滤波,桥接层也拒绝意外的短片段响应。
min_split_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_MIN_SEGMENT_MS", "3000")))
while len(usable) > 1:
short_index = next(
(
index
for index, item in enumerate(usable)
if item["end_time_ms"] - item["start_time_ms"] < min_split_ms
),
None,
)
if short_index is None:
break
if short_index == 0:
usable[1]["start_time_ms"] = usable[0]["start_time_ms"]
del usable[0]
elif short_index == len(usable) - 1:
usable[-2]["end_time_ms"] = usable[-1]["end_time_ms"]
del usable[-1]
else:
previous = usable[short_index - 1]
following = usable[short_index + 1]
if previous["end_time_ms"] - previous["start_time_ms"] >= (
following["end_time_ms"] - following["start_time_ms"]
):
previous["end_time_ms"] = usable[short_index]["end_time_ms"]
del usable[short_index]
else:
following["start_time_ms"] = usable[short_index]["start_time_ms"]
del usable[short_index]
if len(usable) == 1 or len(text) < len(usable):
usable = [max(usable, key=lambda item: item["end_time_ms"] - item["start_time_ms"])]
usable[0]["text"] = text
return usable
total_duration = sum(item["end_time_ms"] - item["start_time_ms"] for item in usable)
char_start = 0
elapsed = 0.0
for index, item in enumerate(usable):
elapsed += item["end_time_ms"] - item["start_time_ms"]
if index == len(usable) - 1:
char_end = len(text)
else:
char_end = round(len(text) * elapsed / total_duration)
char_end = max(char_start + 1, min(char_end, len(text) - (len(usable) - index - 1)))
item["text"] = text[char_start:char_end].strip()
char_start = char_end
return [item for item in usable if item.get("text")]
class BrowserSession: class BrowserSession:
"""Own one browser/native WS pair and translate their message contracts.""" """Own one browser/native WS pair and translate their message contracts."""
@ -190,6 +281,28 @@ class BrowserSession:
} }
await self.emit({"type": "sentences", "sentences": [sentence]}) await self.emit({"type": "sentences", "sentences": [sentence]})
def align_turn_audio_to_vad(self, message: dict[str, Any]) -> None:
"""用原生 VAD 的绝对时间裁掉当前 turn 的前后静音样本。"""
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
try:
speech_start_ms = message.get("speech_start_ms")
if speech_start_ms is not None:
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_start_ms)))
trim_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
trim_bytes -= trim_bytes % 2
del self.turn_audio[:trim_bytes]
self.turn_start_ms += trim_bytes / PCM_BYTES_PER_MS
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
speech_end_ms = message.get("speech_end_ms")
if speech_end_ms is not None:
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_end_ms)))
keep_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
keep_bytes -= keep_bytes % 2
del self.turn_audio[keep_bytes:]
except (TypeError, ValueError):
LOGGER.warning("Ignoring invalid native FunASR VAD boundary: %s", message)
async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None: async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None:
if not frame: if not frame:
return return
@ -236,6 +349,9 @@ class BrowserSession:
if message.get("is_end"): if message.get("is_end"):
self.native_ack = message self.native_ack = message
return return
if message.get("event") == "vad":
self.align_turn_audio_to_vad(message)
continue
text = str(message.get("text") or "") text = str(message.get("text") or "")
if text: if text:
# FunASR online sends the newly decoded text for each chunk. # FunASR online sends the newly decoded text for each chunk.
@ -244,37 +360,29 @@ class BrowserSession:
final_text = self.turn_text.strip() final_text = self.turn_text.strip()
audio = bytes(self.turn_audio) audio = bytes(self.turn_audio)
start_ms = self.turn_start_ms start_ms = self.turn_start_ms
end_ms = self.total_audio_ms end_ms = start_ms + len(audio) / PCM_BYTES_PER_MS
turn_sentence_id = self.sentence_id
if final_text: if final_text:
# Punctuation runs on finalized turns; optional-model failures # 先结束前端 interim 气泡;标点和 CAM++ 在独立 worker 中完成,
# leave the raw ASR text usable. # 不阻塞 native WS 继续读取后续音频帧。
try:
punctuation = await self.auxiliary.punctuate(final_text)
if not punctuation.get("available", False):
LOGGER.warning(
"FunASR punctuation is unavailable: %s",
punctuation.get("error", "model is not configured"),
)
else:
punctuated_text = str(punctuation.get("text") or "").strip()
if punctuated_text:
final_text = punctuated_text
except Exception:
LOGGER.exception("FunASR punctuation request failed; keeping raw text")
await self.emit_sentence( await self.emit_sentence(
final_text, final=True, start_time_ms=start_ms, end_time_ms=end_ms final_text,
final=True,
sentence_id=turn_sentence_id,
start_time_ms=start_ms,
end_time_ms=end_ms,
) )
# CAM++ runs off the native WS reader so ASR keeps draining.
await self.speaker_jobs.put( await self.speaker_jobs.put(
SpeakerJob( SpeakerJob(
sentence_id=self.sentence_id, sentence_id=turn_sentence_id,
text=final_text, text=final_text,
audio=audio, audio=audio,
start_time_ms=start_ms, start_time_ms=start_ms,
end_time_ms=end_ms, end_time_ms=end_ms,
) )
) )
self.sentence_id += 1 # 为同一个 VAD turn 内可能拆出的多个气泡预留独立 ID。
self.sentence_id += TURN_SENTENCE_ID_STRIDE
self.turn_text = "" self.turn_text = ""
self.turn_audio.clear() self.turn_audio.clear()
self.turn_start_ms = self.total_audio_ms self.turn_start_ms = self.total_audio_ms
@ -288,14 +396,52 @@ class BrowserSession:
await self.emit({"type": "error", "message": f"FunASR realtime WS: {exc}"}) await self.emit({"type": "error", "message": f"FunASR realtime WS: {exc}"})
async def resolve_speakers(self) -> None: async def resolve_speakers(self) -> None:
"""Resolve final utterances in order to keep CAM++ cluster IDs stable.""" """按序定稿标点并用 CAM++ 滑窗恢复 turn 内说话人切换。"""
while True: while True:
job = await self.speaker_jobs.get() job = await self.speaker_jobs.get()
try: try:
if job is None: if job is None:
return return
speaker: dict[str, Any] = {"speaker_id": -1} final_text = job.text
try:
punctuation = await self.auxiliary.punctuate(final_text)
if not punctuation.get("available", False):
LOGGER.warning(
"FunASR punctuation is unavailable: %s",
punctuation.get("error", "model is not configured"),
)
else:
punctuated = str(punctuation.get("text") or "").strip()
if punctuated:
final_text = punctuated
except Exception:
LOGGER.exception("FunASR punctuation request failed; keeping raw text")
subsegments: list[dict[str, Any]] = []
if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES: if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
# required CAM++ 按 FunASR 1.5s/0.75s 滑窗识别同一 VAD turn 内的
# 多人切换;若窗长不足或接口暂不可用,再退回整段声纹验证。
track = getattr(self.auxiliary, "track_speakers", None)
if callable(track):
try:
tracked = await track(
job.audio,
self.session_id,
job.start_time_ms,
job.end_time_ms,
)
subsegments = split_text_by_speaker_segments(
final_text,
tracked,
job.start_time_ms,
job.end_time_ms,
)
except Exception:
LOGGER.exception(
"CAM++ sliding-window diarization failed; falling back to whole-turn speaker"
)
if not subsegments and len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
try: try:
resolved = await self.auxiliary.resolve_speaker( resolved = await self.auxiliary.resolve_speaker(
job.audio, job.audio,
@ -303,18 +449,36 @@ class BrowserSession:
job.start_time_ms, job.start_time_ms,
job.end_time_ms, job.end_time_ms,
) )
if resolved: if resolved and int(resolved.get("speaker_id", -1)) >= 0:
speaker = resolved subsegments = [
{
"text": final_text,
"speaker": resolved,
"start_time_ms": job.start_time_ms,
"end_time_ms": job.end_time_ms,
}
]
except Exception: except Exception:
LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id) LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id)
# Re-emit the final sentence with the CAM++ label for the unchanged UI.
if not subsegments:
subsegments = [
{
"text": final_text,
"speaker": {"speaker_id": -1},
"start_time_ms": job.start_time_ms,
"end_time_ms": job.end_time_ms,
}
]
for index, segment in enumerate(subsegments):
await self.emit_sentence( await self.emit_sentence(
job.text, str(segment["text"]),
final=True, final=True,
speaker=speaker, speaker=segment.get("speaker"),
sentence_id=job.sentence_id, sentence_id=job.sentence_id + index,
start_time_ms=job.start_time_ms, start_time_ms=float(segment.get("start_time_ms", job.start_time_ms)),
end_time_ms=job.end_time_ms, end_time_ms=float(segment.get("end_time_ms", job.end_time_ms)),
) )
finally: finally:
self.speaker_jobs.task_done() self.speaker_jobs.task_done()