Improve FunASR realtime speaker tracking
parent
1a239cffbf
commit
14420427a9
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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("<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
|
||||
|
||||
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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue