From 14420427a9d194d86d05e5843bda8f65b74c7edb Mon Sep 17 00:00:00 2001 From: Bifang <915779419@qq.com> Date: Thu, 24 Sep 2026 09:35:55 +0800 Subject: [PATCH] Improve FunASR realtime speaker tracking --- .env.example | 7 + backend/auxiliary_server.py | 407 ++++++++++++++++-- .../realtime_websocket/auxiliary_service.py | 29 ++ .../realtime_websocket/funasr_native_wss.py | 27 +- backend/realtime_websocket/funasr_server.py | 230 ++++++++-- 5 files changed, 622 insertions(+), 78 deletions(-) diff --git a/.env.example b/.env.example index 287172a..7f8b389 100644 --- a/.env.example +++ b/.env.example @@ -13,6 +13,13 @@ FUNASR_ENCODER_LOOK_BACK=4 FUNASR_DECODER_LOOK_BACK=1 FUNASR_VAD_CHUNK_MS=200 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. AUXILIARY_SERVICE_URL=http://127.0.0.1:8010 diff --git a/backend/auxiliary_server.py b/backend/auxiliary_server.py index 3cb959b..d14c0f9 100644 --- a/backend/auxiliary_server.py +++ b/backend/auxiliary_server.py @@ -376,6 +376,332 @@ class AuxiliaryRuntime: output = self._run_embedding_pipeline(model_pipeline, audio_path) return self._normalize_embedding(output) + @staticmethod + def _embedding_rows(result: Any, expected_count: int) -> Any: + """兼容 ModelScope 批量 embedding 的 Tensor、数组和逐条结果格式。""" + import numpy as np + + values = AuxiliaryRuntime._extract_embedding_value(result) + if values is None: + raise RuntimeError("speaker pipeline returned no batch embeddings") + detach = getattr(values, "detach", None) + if callable(detach): + values = detach() + cpu = getattr(values, "cpu", None) + if callable(cpu): + values = cpu() + numpy_method = getattr(values, "numpy", None) + if callable(numpy_method): + values = numpy_method() + try: + 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(" 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 + + embedding = self._normalize_embedding(embedding) + clusters = self.speaker_clusters.setdefault(session_id, []) + best_cluster: dict[str, Any] | None = None + best_score = -1.0 + for cluster in clusters: + if embedding.shape != cluster["embedding"].shape: + raise RuntimeError("speaker embedding dimension changed within the session") + score = float(np.dot(embedding, cluster["embedding"])) + if score > best_score: + best_score = score + best_cluster = cluster + + if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: + count = int(best_cluster["count"]) + best_cluster["embedding"] = self._normalize_embedding( + (best_cluster["embedding"] * count) + embedding + ) + best_cluster["count"] = count + 1 + speaker_id = int(best_cluster["speaker_id"]) + confidence = best_score + strategy = "online_embedding_cluster_match" + else: + # 不设置会话人数上限;与已有中心不匹配的声纹建立新身份。 + speaker_id = len(clusters) + clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1}) + confidence = 0.75 + strategy = "online_embedding_cluster_new" + + return { + "speaker_id": speaker_id, + "speaker_name": f"说话人 {speaker_id + 1}", + "speaker_evidence": "fresh", + "speaker_confidence": round(max(0.6, min(1.0, confidence)), 3), + "speaker_strategy": strategy, + "speaker_status": "confirmed", + "speaker_reason": "CAM++ 滑窗声纹已完成会话内在线聚类", + "start_time": start_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, @@ -395,49 +721,12 @@ class AuxiliaryRuntime: if embedding is None: return {"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0, "speaker_status": "insufficient_audio", "speaker_reason": "音频不足 800ms,未提取声纹"} - import numpy as np - # reset 可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。 if session_id not in self.speaker_last_seen: return None - embedding = self._normalize_embedding(embedding) - clusters = self.speaker_clusters.setdefault(session_id, []) - best_cluster: dict[str, Any] | None = None - best_score = -1.0 - for cluster in clusters: - if embedding.shape != cluster["embedding"].shape: - raise RuntimeError("speaker embedding dimension changed within the session") - score = float(np.dot(embedding, cluster["embedding"])) - if score > best_score: - best_score = score - best_cluster = cluster - - if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD: - count = int(best_cluster["count"]) - best_cluster["embedding"] = self._normalize_embedding( - (best_cluster["embedding"] * count) + embedding - ) - best_cluster["count"] = count + 1 - speaker_id = int(best_cluster["speaker_id"]) - confidence = best_score - strategy = "online_embedding_cluster_match" - else: - speaker_id = len(clusters) - clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1}) - confidence = 0.75 - strategy = "online_embedding_cluster_new" - - return { - "speaker_id": speaker_id, - "speaker_name": f"说话人 {speaker_id + 1}", - "speaker_evidence": "fresh", - "speaker_confidence": round(max(0.6, min(1.0, confidence)), 3), - "speaker_strategy": strategy, - "speaker_status": "confirmed", - "speaker_reason": "当前片段独立声纹已完成在线聚类", - "start_time": start_time_ms, - "end_time": end_time_ms, - } + return self._assign_embedding( + session_id, embedding, start_time_ms, end_time_ms + ) def reset_speaker_session(self, session_id: str) -> None: """释放已结束 WebSocket 的聚类中心,防止长时间运行时内存增长。""" @@ -627,6 +916,43 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response: 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: """释放已经结束的实时会话聚类中心。""" 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/diarization", diarization_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) return app diff --git a/backend/realtime_websocket/auxiliary_service.py b/backend/realtime_websocket/auxiliary_service.py index bbe4929..065ee68 100644 --- a/backend/realtime_websocket/auxiliary_service.py +++ b/backend/realtime_websocket/auxiliary_service.py @@ -116,6 +116,35 @@ class AuxiliaryModelService: # 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。 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: """通知辅助服务释放当前 WebSocket 对应的在线聚类状态。""" if self._session is None: diff --git a/backend/realtime_websocket/funasr_native_wss.py b/backend/realtime_websocket/funasr_native_wss.py index db0bcdb..efb52b1 100644 --- a/backend/realtime_websocket/funasr_native_wss.py +++ b/backend/realtime_websocket/funasr_native_wss.py @@ -650,6 +650,20 @@ async def ws_serve(websocket, path=None): record_error(f"vad inference failed: {e}") 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: speech_start = True 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): 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 message = { "mode": mode, - "text": rec_result["text"], + "text": rec_result.get("text", ""), "wav_name": websocket.wav_name, - "is_final": bool( - websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking) - ), + "is_final": is_final, } await websocket.send(json.dumps(message, ensure_ascii=False)) diff --git a/backend/realtime_websocket/funasr_server.py b/backend/realtime_websocket/funasr_server.py index 67e22b9..d56ba52 100644 --- a/backend/realtime_websocket/funasr_server.py +++ b/backend/realtime_websocket/funasr_server.py @@ -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 MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2) 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) SESSION_REGISTRY_KEY = web.AppKey("sessions", dict) @@ -121,6 +122,96 @@ class SpeakerJob: 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: """Own one browser/native WS pair and translate their message contracts.""" @@ -190,6 +281,28 @@ class BrowserSession: } 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: if not frame: return @@ -236,6 +349,9 @@ class BrowserSession: if message.get("is_end"): self.native_ack = message return + if message.get("event") == "vad": + self.align_turn_audio_to_vad(message) + continue text = str(message.get("text") or "") if text: # FunASR online sends the newly decoded text for each chunk. @@ -244,37 +360,29 @@ class BrowserSession: final_text = self.turn_text.strip() audio = bytes(self.turn_audio) 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: - # Punctuation runs on finalized turns; optional-model failures - # leave the raw ASR text usable. - 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") + # 先结束前端 interim 气泡;标点和 CAM++ 在独立 worker 中完成, + # 不阻塞 native WS 继续读取后续音频帧。 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( SpeakerJob( - sentence_id=self.sentence_id, + sentence_id=turn_sentence_id, text=final_text, audio=audio, start_time_ms=start_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_audio.clear() 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}"}) async def resolve_speakers(self) -> None: - """Resolve final utterances in order to keep CAM++ cluster IDs stable.""" + """按序定稿标点并用 CAM++ 滑窗恢复 turn 内说话人切换。""" while True: job = await self.speaker_jobs.get() try: if job is None: 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: + # 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: resolved = await self.auxiliary.resolve_speaker( job.audio, @@ -303,19 +449,37 @@ class BrowserSession: job.start_time_ms, job.end_time_ms, ) - if resolved: - speaker = resolved + if resolved and int(resolved.get("speaker_id", -1)) >= 0: + subsegments = [ + { + "text": final_text, + "speaker": resolved, + "start_time_ms": job.start_time_ms, + "end_time_ms": job.end_time_ms, + } + ] except Exception: 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. - await self.emit_sentence( - job.text, - final=True, - speaker=speaker, - sentence_id=job.sentence_id, - start_time_ms=job.start_time_ms, - end_time_ms=job.end_time_ms, - ) + + 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( + str(segment["text"]), + final=True, + speaker=segment.get("speaker"), + sentence_id=job.sentence_id + index, + start_time_ms=float(segment.get("start_time_ms", job.start_time_ms)), + end_time_ms=float(segment.get("end_time_ms", job.end_time_ms)), + ) finally: self.speaker_jobs.task_done()