omp修复版本

Bifang 2026-09-10 14:23:18 +08:00
parent 6243bb30f2
commit c76e70d10e
8 changed files with 513 additions and 112 deletions

2
.gitignore vendored 100644
View File

@ -0,0 +1,2 @@
__pycache__/
.env

View File

@ -1,87 +0,0 @@
# Demo 说话人链路修复与验收
本次仅修改 `demo/`。参考原 `app/services/qwen3_websocket_asr.py` 的前滚音频、同句异步说话人更新和结束提交方式,保留独立 vLLM ASR + 辅助声纹服务 + WebSocket 三个进程。当前机器没有显卡,验证采用模拟模型,真实声纹准确率和模型加载状态需在部署机器上验收。
## 已确认的代码问题
1. 页面只处理 `sentences`,忽略已实现的 `display_state`,因此后端合并结果没有展示。未知片段还会临时附在上一位已确认说话人的气泡里。
2. 页面点击停止后 5 秒强制断开,但辅助 HTTP 请求的超时是 45 秒,可能截掉迟到的声纹更新。
3. 无 embedding、低置信度、协议字段不完整等情况被静默丢弃用户无法区分等待、短音频、服务错误和证据拒绝。
4. 旧声纹长度检查把句尾 800ms 静音一起计算可能让短插话通过长度要求WAV 固定去掉 44 字节也会破坏带额外元数据的输入。
5. 仅 vLLM 启动器加载 `demo/.env`,另两个进程没有加载;`0.6b` 下载别名也可能被直接当作公开模型名发送。
这些是代码中可复现的问题,不能据此断言部署中的声纹模型一定已经正常加载。新版本将具体原因直接显示在片段气泡和日志里。
## 修复后的行为
- partial/final 和声纹更新继续使用同一 `sentence_id`。页面显示完整聚合快照,过时的 `revision` 不覆盖新快照。
- 仅合并相邻且身份可信的片段。未知短插话独立显示A→B→A 保留顺序。不同实名或实名与弱匿名身份不会因为簇 ID 相同而合并。
- 每条声纹证据必须来自当前片段;同步、异步两种入口都拒绝 `short_attach` / `embedding_attach`。不复制上一段 embedding。声纹向量拒绝零向量、NaN、Infinity 和多样本矩阵。
- WebSocket 以有效有声帧判断长度。少于 800ms 的语音保持 pending800ms1.6s 也必须独立提取特征,不能直接继承前一位身份。该长度门槛属于保守保护,不能保证短样本识别准确率。
- 保留 200ms 前滚,提交时去除尾部静音。增量解析 WAV 的 RIFF/fmt/data/JUNK 等头,拒绝非 16kHz 单声道 PCM16 文件。
- `stop/eof` 返回 `draining`,排空 ASR 和声纹队列后才发送最终快照及 `end`。声纹失败不阻断 ASR。abort、断线、异常和正常结束均清理聚类会话。
- ASR worker 异常会立即发送 `error`,不会一直等待客户端停止。
## 如何判断卡在哪里
页面顶部显示已确认数量;未匹配到说话人的气泡统一显示“未知说话人”,详细等待或失败原因通过标签悬停提示和原始日志查看。原始日志保留 `speaker_status`、`speaker_reason`、`speaker_strategy` 和置信度。
| `speaker_status` | 含义与排查方向 |
|---|---|
| `waiting_final` | 讲话仍在进行,等待静音或最大时长切段 |
| `queued` / `processing` | 文本已完成,声纹正在排队或推理 |
| `confirmed` | 当前片段声纹已确认,应该显示说话人标签 |
| `insufficient_audio` | 有效语音过短,不继承上一位;使用较长发言复测 |
| `service_unavailable` / `service_error` | 未配置、无法访问或模型推理失败;检查辅助服务及完整错误 |
| `no_embedding` | 服务没有产生可用特征 |
| `evidence_rejected` | 缺少 fresh/confirmed 证据、置信度不足或使用了继承策略 |
| `disabled` | 本次未开启说话人分离 |
健康检查为 `http://辅助服务器:8010/health`。新版本有 `speaker_protocol_version: 2`,重点检查 `speaker_embedding_ready``speaker_embedding_model`。若没有版本字段检查是否重启了更新后的辅助进程。HTTP 健康检查不执行真实声纹推理,不能代替音频验收。
辅助服务目前返回匿名的“说话人 1、2……”demo 没有接入原应用的声纹注册库因此不会自动识别人员实名。vLLM 仅输出转写文本,声纹标签由 `scripts/auxiliary_server.py` 负责。
模型职责要区分:`iic/speech_campplus_sv_zh-cn_16k-common` 是实时 turn 的 CAM++ embedding 模型,必须加载;`iic/speech_campplus_speaker-diarization_common` 是完整音频分离 pipeline包含额外的 change locator/VAD 依赖,当前实时 WebSocket 不在启动阶段调用它。WebSocket 自身仍用轻量 RMS 帧门控切句,辅助服务的 FunASR VAD 对外提供 `/v1/vad`,并供完整 diarization 依赖使用;因此启动辅助服务是为了 CAM++ 声纹和聚类,不能把整段 diarization 的加载失败误认为 vLLM 失败。
## 部署后操作
在部署机更新这些文件后,已有 vLLM 可继续运行。重新安装增补的 `python-dotenv` 依赖,并重启辅助服务及 WebSocket。辅助服务启动时强制依赖 VAD + CAM++ `speaker_verification`;完整 CAM++ diarization 不再阻断实时启动;以下命令都从 `demo/` 目录执行:
```text
pip install -r requirements-auxiliary.txt
pip install -r realtime_asr_optimization_demo/requirements.txt
python scripts/auxiliary_server.py
```
如果仍提示核心模型缺失或加载失败,日志会列出模型 ID、实际查找路径、状态和底层异常。先执行 `python scripts/download_models.py --auxiliary-only`,或设置 `.env``MODEL_DIR` 指向同时包含 `damo/speech_fsmn_vad_zh-cn-16k-common-pytorch``iic/speech_campplus_sv_zh-cn_16k-common` 的目录。不要用 `AUXILIARY_ALLOW_MISSING=true` 掩盖 VAD/CAM++ 核心模型缺失;该选项只适合临时查看可选模型状态。
另一个终端执行:
```text
python realtime_asr_optimization_demo/server.py --no-browser
```
三个进程现在均读取 `demo/.env`,系统环境变量优先。确认 `MODEL_SERVICE_URL` 指向现有 vLLM 的 `/v1``AUXILIARY_SERVICE_URL` 指向辅助服务。`QWEN3_ASR_MODEL` 支持部署清单别名,设置 `VLLM_SERVED_MODEL_NAME` 时优先使用该公开名称。更新后刷新浏览器;脚本 URL 已更新版本号。
`127.0.0.1` 指各 Python 服务运行的机器,不是浏览器所在机器。跨服务器部署时填写对应服务器 IP。浏览器麦克风访问远程页面需要安全上下文HTTPS本机 localhost 可用于测试。
## 验收顺序
1. 单人讲话 24 秒后停顿:先出现文本,随后同句变为“说话人 1”。
2. 同一人再次讲话并停顿:确认后相邻块应合并;取消页面“合并相邻”可对照物理片段。
3. A→B→A各说 2 秒以上:应保持三个时间顺序块。标签准确率需要真实声纹模型验证。
4. A 后 B 说一个很短的“嗯”:应显示“有效语音不足”,不能进入 A 的气泡。
5. 辅助服务关闭时测试ASR 仍完成,片段明确显示服务错误。
6. 讲话中点击停止:等待最终结果,不能在五秒时丢失说话人更新。
无显卡回归命令:
```text
cd demo
python -m unittest discover -s tests -v
cd realtime_asr_optimization_demo
python -m unittest discover -s tests -v
node --test tests/test_frontend.cjs
```
这些测试覆盖状态机、模拟 HTTP/WebSocket、延迟更新、前端脚本和有效向量校验不执行模型下载或 GPU 推理。当前 VLLM 适配器仍是 HTTP 累积窗口 partial`native_partial_supported=false`;本次没有把 HTTP 接口包装成原生增量模型状态。单个内部片段中无停顿的多人换话或重叠讲话仍需真实模型和更细粒度切段验证。

View File

@ -86,7 +86,7 @@ python server.py --no-browser
1. 选择 Mic 或 File。 1. 选择 Mic 或 File。
2. 点击开始,浏览器通过 `/ws` 建立本项目 WebSocket。 2. 点击开始,浏览器通过 `/ws` 建立本项目 WebSocket。
3. 页面展示腾讯 Demo 风格的气泡;麦克风或 PCM/WAV 文件按流式方式输入partial 会在讲话过程中实时刷新。 3. 页面展示腾讯 Demo 风格的气泡;麦克风或 PCM/WAV 文件按流式方式输入partial 会在讲话过程中实时刷新。
4. VAD 检测到静音提交当前 turn先返回 pending再异步更新说话人。 4. VAD 检测到静音、或活跃 turn 内声纹窗口比对检出换人时提交当前 turn先返回 pending再异步更新说话人。
5. 停止会发送 `stop`,服务端完成当前 turn 和 speaker 队列后再发送 `end` 5. 停止会发送 `stop`,服务端完成当前 turn 和 speaker 队列后再发送 `end`
6. `abort` 只取消会话,不提交当前片段。 6. `abort` 只取消会话,不提交当前片段。
@ -105,13 +105,23 @@ python server.py --no-browser
"enable_native_partial_stream": true, "enable_native_partial_stream": true,
"partial_interval_ms": 1200, "partial_interval_ms": 1200,
"max_segment_sec": 12, "max_segment_sec": 12,
"display_merge": true "display_merge": true,
"speaker_window_ms": 1000,
"speaker_switch_similarity": 0.5
} }
``` ```
`speaker_window_ms``speaker_switch_similarity` 可省略,默认读取环境变量
`SPEAKER_TURN_WINDOW_MS`1000`SPEAKER_SWITCH_SIMILARITY`0.5)。
随后持续发送 16kHz、单声道、PCM16 二进制音频;文件模式仅支持 PCM/WAV结束发送 `{"type":"eof"}`,停止发送 随后持续发送 16kHz、单声道、PCM16 二进制音频;文件模式仅支持 PCM/WAV结束发送 `{"type":"eof"}`,停止发送
`{"type":"stop"}`,取消发送 `{"type":"abort"}`。每个已完成 turn 会向辅助服务 `{"type":"stop"}`,取消发送 `{"type":"abort"}`。每个已完成 turn 会向辅助服务
发送一次 `/v1/speaker/resolve`,只包含当前 turn 音频和 session_id不会重复上传整段会话。 发送一次 `/v1/speaker/resolve`,只包含当前 turn 音频和 session_id不会重复上传整段会话。
换人强制切段产生的边界段会附带 `speaker_verified=0`:辅助服务仍返回匹配标签,
但不更新簇质心、不建立新簇。活跃 turn 内另有窗口比对请求 `/v1/speaker/embedding`
multipart 音频,返回 `{"ok": true, "embedding": [...]}`),只出向量、零聚类副作用。
辅助服务 `speaker_protocol_version` 为 3旧版辅助服务缺少窗端点时 WebSocket 侧
探测一次即自动停用切换感知切段并发送 `speaker_warning`,不影响 ASR。
服务端会发送 `start`、`sentences`、`display_state`、`metrics`、`speaker_warning`、`draining`、`end` 和 `error` 服务端会发送 `start`、`sentences`、`display_state`、`metrics`、`speaker_warning`、`draining`、`end` 和 `error`
页面用带 `revision``display_state` 渲染,以 `block_id` 标识展示块;`sentences` 保留原始片段及诊断状态。 页面用带 `revision``display_state` 渲染,以 `block_id` 标识展示块;`sentences` 保留原始片段及诊断状态。
@ -119,6 +129,10 @@ python server.py --no-browser
`sentences``sentence_type=0` 是 partial`sentence_type=1` 是 final同一个 `sentences``sentence_type=0` 是 partial`sentence_type=1` 是 final同一个
`sentence_id` 必须覆盖更新而不是追加。`sentence_strategy=0` 使用约 800ms `sentence_id` 必须覆盖更新而不是追加。`sentence_strategy=0` 使用约 800ms
静音切句,`sentence_strategy=1` 使用约 1400ms 静音切段,更适合段落模式。 静音切句,`sentence_strategy=1` 使用约 1400ms 静音切段,更适合段落模式。
静音阈值之外还有说话人切换感知切段:活跃 turn 内每积累约 1 秒有声窗口就与
辅助服务比对一次声纹,与当前 turn 参考向量余弦低于阈值时在“停顿≥200ms 后的
首个起音”处强制切开,头段 final 的 `commit_reason``speaker_change`。这覆盖
上一位尾句与下一位间隔不足静音阈值(常见 200~600ms 换话停顿)被合成一句的问题。
## 目录 ## 目录

View File

@ -36,6 +36,8 @@ class AuxiliaryModelService:
def __init__(self, config: AuxiliaryServiceConfig) -> None: def __init__(self, config: AuxiliaryServiceConfig) -> None:
self.config = config self.config = config
self._session: ClientSession | None = None self._session: ClientSession | None = None
# 旧版辅助服务没有窗口声纹端点;探测到一次 404/405 后不再重复请求。
self.speaker_embedding_unsupported = False
async def start(self) -> None: async def start(self) -> None:
"""创建可复用的 HTTP 会话,避免每个片段重复建立 TCP 连接。""" """创建可复用的 HTTP 会话,避免每个片段重复建立 TCP 连接。"""
@ -70,11 +72,14 @@ class AuxiliaryModelService:
session_id: str, session_id: str,
start_time_ms: float, start_time_ms: float,
end_time_ms: float, end_time_ms: float,
speaker_verified: bool = True,
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""提交一个已经由实时 VAD 完成的 turn获取在线聚类结果。 """提交一个已经由实时 VAD 完成的 turn获取在线聚类结果。
每次请求只包含当前 turn不上传整段会话辅助服务通过 session_id 每次请求只包含当前 turn不上传整段会话辅助服务通过 session_id
保存聚类中心因此同一说话人在 ABA 场景下仍能保持同一标签 保存聚类中心因此同一说话人在 ABA 场景下仍能保持同一标签
speaker_verified=False 表示换人强制切段产生的边界段辅助服务
只匹配标签不用它更新簇质心
""" """
if self._session is None: if self._session is None:
raise RuntimeError("auxiliary model service is not started") raise RuntimeError("auxiliary model service is not started")
@ -83,6 +88,7 @@ class AuxiliaryModelService:
form.add_field("session_id", session_id) form.add_field("session_id", session_id)
form.add_field("start_time_ms", str(start_time_ms)) form.add_field("start_time_ms", str(start_time_ms))
form.add_field("end_time_ms", str(end_time_ms)) form.add_field("end_time_ms", str(end_time_ms))
form.add_field("speaker_verified", "1" if speaker_verified else "0")
endpoint = self.config.base_url.rstrip("/") + "/v1/speaker/resolve" endpoint = self.config.base_url.rstrip("/") + "/v1/speaker/resolve"
async with self._session.post(endpoint, data=form) as response: async with self._session.post(endpoint, data=form) as response:
body = await response.text() body = await response.text()
@ -99,6 +105,36 @@ class AuxiliaryModelService:
# 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。 # 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。
return decoded return decoded
async def speaker_embedding(self, pcm_bytes: bytes) -> list[float] | None:
"""为一个活跃 turn 的短窗口提取归一化声纹,绝不读取或更新聚类状态。
任何失败都返回 None 而不是抛异常窗比对只是切段辅助绝不能
把辅助服务的抖动传导成 ASR 阻塞404/405 视为服务版本过旧置位
speaker_embedding_unsupported WebSocket 侧停用该功能"""
if self._session is None or self.speaker_embedding_unsupported:
return None
form = FormData()
form.add_field("file", pcm16_to_wav(pcm_bytes), filename="window.wav", content_type="audio/wav")
endpoint = self.config.base_url.rstrip("/") + "/v1/speaker/embedding"
try:
# 窗比对嵌在音频帧循环里,必须用远短于会话级 45s 的超时兜底。
async with self._session.post(endpoint, data=form, timeout=ClientTimeout(total=8.0)) as response:
if response.status in {404, 405}:
self.speaker_embedding_unsupported = True
return None
if response.status >= 400:
return None
decoded = await response.json(content_type=None)
except Exception:
return None
embedding = decoded.get("embedding") if isinstance(decoded, dict) else None
if not isinstance(embedding, list) or not embedding:
return None
try:
return [float(value) for value in embedding]
except (TypeError, ValueError):
return None
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

@ -41,9 +41,30 @@ VAD_SILENCE_MS = 800
PARAGRAPH_SILENCE_MS = 1400 PARAGRAPH_SILENCE_MS = 1400
VAD_RMS_THRESHOLD = 450 VAD_RMS_THRESHOLD = 450
MIN_SPEAKER_VOICE_MS = 800 MIN_SPEAKER_VOICE_MS = 800
# 说话人切换感知切段:活跃 turn 内每积累一个窗口就用声纹比对一次,
# 与当前 turn 参考向量的余弦低于阈值时视为换人,在窗口边界强制切段。
# 静音阈值800/1400ms覆盖不了常见的 200~600ms 换话停顿,窗比对是它的补充。
SPEAKER_TURN_WINDOW_MS = max(800, int(os.getenv("SPEAKER_TURN_WINDOW_MS", "1000")))
SPEAKER_SWITCH_SIMILARITY = float(os.getenv("SPEAKER_SWITCH_SIMILARITY", "0.5"))
SPEAKER_WINDOW_MAX_FAILURES = 3
# 窗口里的有声时长不足时比对没有意义(长静音窗声纹不可信),攒够再查,
# 同时避免静音段每帧重复触发请求。
SPEAKER_WINDOW_MIN_VOICED_MS = 600
# 失配窗内定位精确切换点:换话几乎总有停顿,取首个"停顿≥该时长后的起音"
# 作为切点,避免把上一位的话尾按窗口起点粗暴划给下一位。
SPEAKER_SWITCH_GAP_MS = 200
LOGGER = logging.getLogger(__name__) LOGGER = logging.getLogger(__name__)
def _coerce_similarity(value: Any, default: float) -> float:
"""解析 start 消息里的相似度阈值,非法值回退到默认,避免 NaN 进比较。"""
try:
parsed = float(value)
except (TypeError, ValueError):
return default
return parsed if math.isfinite(parsed) and -1.0 <= parsed <= 1.0 else default
class EndOfStream: class EndOfStream:
"""带明确类型的队列结束标记,用于区分控制信号和真实音频字节。""" """带明确类型的队列结束标记,用于区分控制信号和真实音频字节。"""
@ -62,6 +83,9 @@ class SpeakerJob:
start_time_ms: float start_time_ms: float
end_time_ms: float end_time_ms: float
voiced_ms: float = 0.0 voiced_ms: float = 0.0
# 强制切段产生的头段可能混入下一说话人至多一个窗口的音频,标记后
# 辅助服务只用它匹配标签、不更新簇质心,避免污染在线聚类。
speaker_verified: bool = True
def validate_model_service_url(value: str) -> str: def validate_model_service_url(value: str) -> str:
@ -141,6 +165,19 @@ class RealtimeSession:
self.in_speech = False self.in_speech = False
self.voiced_ms = 0.0 self.voiced_ms = 0.0
self.pre_roll = bytearray() self.pre_roll = bytearray()
# 窗级说话人切换检测状态:参考向量是当前段已通过比对窗口的均值,
# window_* 记录自上次比对以来累积的音频字节与有声时长。
self.turn_window_ms = max(800, int(start.get("speaker_window_ms") or SPEAKER_TURN_WINDOW_MS))
self.switch_similarity = _coerce_similarity(start.get("speaker_switch_similarity"), SPEAKER_SWITCH_SIMILARITY)
self.turn_window_bytes = int(self.turn_window_ms / 1000 * PARTIAL_BYTES_PER_SECOND)
self.turn_reference: list[float] | None = None
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
self.window_check_failures = 0
self.window_check_available = (
auxiliary_service is not None
and hasattr(auxiliary_service, "speaker_embedding")
)
self.wav_header_buffer = bytearray() self.wav_header_buffer = bytearray()
self.wav_payload_started = self.source != "file" or Path(self.file_name).suffix.lower() != ".wav" self.wav_payload_started = self.source != "file" or Path(self.file_name).suffix.lower() != ".wav"
self.wav_riff_read = False self.wav_riff_read = False
@ -148,6 +185,7 @@ class RealtimeSession:
self.wav_data_remaining: int | None = None self.wav_data_remaining: int | None = None
self.speaker_warning_sent = False self.speaker_warning_sent = False
self.speaker_enabled = self._parse_flag(start.get("speaker_diarization"), True) self.speaker_enabled = self._parse_flag(start.get("speaker_diarization"), True)
self.window_check_available = self.window_check_available and self.speaker_enabled
self.input_stopped = False self.input_stopped = False
@staticmethod @staticmethod
@ -197,10 +235,10 @@ class RealtimeSession:
"""根据 16 kHz PCM 字节数计算时长,不依赖前端可能漂移的时间戳。""" """根据 16 kHz PCM 字节数计算时长,不依赖前端可能漂移的时间戳。"""
return len(self.segment_audio) / PARTIAL_BYTES_PER_SECOND * 1000 return len(self.segment_audio) / PARTIAL_BYTES_PER_SECOND * 1000
async def _transcribe(self, partial: bool) -> str | None: async def _transcribe(self, partial: bool, audio: bytes | None = None) -> str | None:
"""通过 VLLM 适配器转写当前逻辑片段,并保留中间/最终请求的统一入口""" """通过 VLLM 适配器转写当前逻辑片段;换人切段时转写显式传入的头音频"""
return await self.model_service.transcribe( return await self.model_service.transcribe(
bytes(self.segment_audio), bytes(self.segment_audio) if audio is None else audio,
"mic", "mic",
"turn.pcm", "turn.pcm",
partial=partial, partial=partial,
@ -320,6 +358,7 @@ class RealtimeSession:
self.session_id, self.session_id,
job.start_time_ms, job.start_time_ms,
job.end_time_ms, job.end_time_ms,
job.speaker_verified,
) )
except Exception as exc: except Exception as exc:
# 辅助服务异常不能阻断 ASR当前片段继续保持 pending方便定位服务问题。 # 辅助服务异常不能阻断 ASR当前片段继续保持 pending方便定位服务问题。
@ -381,6 +420,124 @@ class RealtimeSession:
self.silence_ms = 0 self.silence_ms = 0
self.in_speech = False self.in_speech = False
self.voiced_ms = 0.0 self.voiced_ms = 0.0
self._reset_turn_embedding_state()
def _reset_turn_embedding_state(self) -> None:
"""段与段之间不共享声纹参考;新 turn 从第一个窗口重新建立参考。"""
self.turn_reference = None
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
@staticmethod
def _unit_average(first: list[float], second: list[float]) -> list[float]:
"""两个单位向量的等权平均并重新归一化,作为滚动参考向量。"""
merged = [a + b for a, b in zip(first, second)]
norm = math.sqrt(sum(value * value for value in merged))
if not math.isfinite(norm) or norm < 1e-8:
return list(first)
return [value / norm for value in merged]
async def _check_speaker_window(self) -> bool:
"""比对活跃段尾部窗口的声纹,检出换人时在窗口边界强制切段。
窗提取失败绝不阻断 ASR连续失败或辅助服务不支持时整体停用
与既有 speaker 失败不阻断转写的策略一致返回是否发生了切段"""
check_len = len(self.segment_audio)
window_audio = bytes(self.segment_audio[self.window_start_bytes:check_len])
embedding = await self.auxiliary_service.speaker_embedding(window_audio)
if embedding is None:
if getattr(self.auxiliary_service, "speaker_embedding_unsupported", False):
self.window_check_available = False
await self.warn_speaker("辅助服务不支持窗口声纹提取,已停用说话人切换感知切段")
else:
self.window_check_failures += 1
if self.window_check_failures >= SPEAKER_WINDOW_MAX_FAILURES:
self.window_check_available = False
await self.warn_speaker("窗口声纹提取连续失败,本次会话已停用切换感知切段")
return False
self.window_check_failures = 0
reference = self.turn_reference
if reference is None or len(reference) != len(embedding):
self.turn_reference = embedding
self.window_start_bytes = check_len
self.window_voiced_ms = 0.0
return False
score = sum(a * b for a, b in zip(reference, embedding))
if score >= self.switch_similarity:
self.turn_reference = self._unit_average(reference, embedding)
self.window_start_bytes = check_len
self.window_voiced_ms = 0.0
return False
split_ms = self.window_start_bytes / PARTIAL_BYTES_PER_SECOND * 1000
LOGGER.info(
"speaker switch detected: session=%s sentence=%s score=%.3f threshold=%.2f at_ms=%.0f",
self.session_id, self.segment_id, score, self.switch_similarity,
self.segment_start_ms + split_ms,
)
await self._split_on_speaker_change(self.window_start_bytes, embedding)
return True
def _find_switch_boundary(self, window_start: int) -> tuple[int, float]:
"""在失配窗内定位精确切换点:首个"停顿≥SPEAKER_SWITCH_GAP_MS 后的起音"
找不到明显停顿就保守回退到窗口起点返回切点字节偏移与切点之后的
有声时长供头尾两段正确归属 voiced_ms"""
gap_frames = max(1, SPEAKER_SWITCH_GAP_MS // VAD_FRAME_MS)
offsets = range(window_start, len(self.segment_audio) - VAD_FRAME_BYTES + 1, VAD_FRAME_BYTES)
voiced_flags = [
self._is_voice_frame(bytes(self.segment_audio[offset:offset + VAD_FRAME_BYTES]))
for offset in offsets
]
boundary = window_start
silence_run = 0
for offset, voiced in zip(offsets, voiced_flags):
if voiced and silence_run >= gap_frames:
boundary = offset
break
silence_run = 0 if voiced else silence_run + 1
tail_voiced_ms = sum(
VAD_FRAME_MS for offset, voiced in zip(offsets, voiced_flags)
if voiced and offset >= boundary
)
return boundary, float(tail_voiced_ms)
async def _split_on_speaker_change(self, window_start: int, new_reference: list[float]) -> None:
"""窗比对检出不同说话人:头段提交 final 并排队声纹,尾段直接成为新 turn。
_commit_segment 不同切点位于活跃语音内部不做尾部静音裁剪
尾段保留全部已收音频并沿用被判为新说话人的窗口向量作为新参考"""
split_at, tail_voiced_ms = self._find_switch_boundary(window_start)
head_audio = bytes(self.segment_audio[:split_at])
head_duration_ms = split_at / PARTIAL_BYTES_PER_SECOND * 1000
head_start_ms = self.segment_start_ms
head_end_ms = head_start_ms + head_duration_ms
head_sentence_id = self.segment_id
head_voiced_ms = max(0.0, self.voiced_ms - tail_voiced_ms)
text = await self._transcribe(partial=False, audio=head_audio)
if text:
await self._emit_transcription(text, 1, head_end_ms, "speaker_change")
if self.speaker_enabled:
await self.speaker_queue.put(
SpeakerJob(
sentence_id=head_sentence_id,
audio=head_audio,
start_time_ms=head_start_ms,
end_time_ms=head_end_ms,
voiced_ms=head_voiced_ms,
speaker_verified=False,
)
)
elif self.assembler.segments.pop(head_sentence_id, None) is not None:
# 与静音提交一致final 无文本时撤回该 id 的 partial不留下悬空等待。
await self.emit_state()
del self.segment_audio[:split_at]
self.segment_start_ms = head_end_ms
self.segment_id += 1
self.silence_ms = 0
self.voiced_ms = tail_voiced_ms
self.turn_reference = new_reference
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
async def process_audio(self) -> None: async def process_audio(self) -> None:
"""消费音频,以 VAD 静音结束作为切句主逻辑,并按窗口发送 partial。""" """消费音频,以 VAD 静音结束作为切句主逻辑,并按窗口发送 partial。"""
@ -412,16 +569,24 @@ class RealtimeSession:
self.segment_audio = bytearray(self.pre_roll) self.segment_audio = bytearray(self.pre_roll)
self.pre_roll.clear() self.pre_roll.clear()
last_partial_bytes = 0 last_partial_bytes = 0
self._reset_turn_embedding_state()
if self.in_speech: if self.in_speech:
self.segment_audio.extend(frame) self.segment_audio.extend(frame)
if voiced: if voiced:
self.voiced_ms += VAD_FRAME_MS self.voiced_ms += VAD_FRAME_MS
self.window_voiced_ms += VAD_FRAME_MS
self.silence_ms = 0 if voiced else self.silence_ms + VAD_FRAME_MS self.silence_ms = 0 if voiced else self.silence_ms + VAD_FRAME_MS
if len(self.segment_audio) - last_partial_bytes >= partial_bytes and self.silence_ms < self.silence_limit_ms: if len(self.segment_audio) - last_partial_bytes >= partial_bytes and self.silence_ms < self.silence_limit_ms:
text = await self._transcribe(partial=True) text = await self._transcribe(partial=True)
if text: if text:
await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms()) await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms())
last_partial_bytes = len(self.segment_audio) last_partial_bytes = len(self.segment_audio)
if (self.window_check_available
and len(self.segment_audio) - self.window_start_bytes >= self.turn_window_bytes
and self.window_voiced_ms >= SPEAKER_WINDOW_MIN_VOICED_MS):
if await self._check_speaker_window():
# 尾段是全新文本,下一帧就允许触发它的 partial。
last_partial_bytes = 0
if self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000: if self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000:
await self._commit_segment("silence" if self.silence_ms >= self.silence_limit_ms else "max_duration") await self._commit_segment("silence" if self.silence_ms >= self.silence_limit_ms else "max_duration")
last_partial_bytes = 0 last_partial_bytes = 0

View File

@ -5,6 +5,9 @@ from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
import asyncio import asyncio
import io import io
import math
import struct
import unittest
import wave import wave
import os import os
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
@ -14,7 +17,9 @@ from aiohttp.test_utils import AioHTTPTestCase
from server import ( from server import (
AUXILIARY_SERVICE_KEY, AUXILIARY_SERVICE_KEY,
EOF,
MODEL_SERVICE_KEY, MODEL_SERVICE_KEY,
RealtimeSession,
config_handler, config_handler,
deployment_model_name, deployment_model_name,
validate_model_service_url, validate_model_service_url,
@ -46,8 +51,9 @@ class FakeAuxiliaryService:
session_id: str, session_id: str,
start_time_ms: float, start_time_ms: float,
end_time_ms: float, end_time_ms: float,
speaker_verified: bool = True,
) -> dict[str, object]: ) -> dict[str, object]:
_ = (audio_bytes, session_id, start_time_ms, end_time_ms) _ = (audio_bytes, session_id, start_time_ms, end_time_ms, speaker_verified)
speaker_id = [0, 1, 0][min(self.calls, 2)] speaker_id = [0, 1, 0][min(self.calls, 2)]
self.calls += 1 self.calls += 1
return { return {
@ -73,12 +79,12 @@ class WebSocketFlowTests(AioHTTPTestCase):
async def collect(self, ws, until="end"): async def collect(self, ws, until="end"):
"""限定等待时间,回归测试中的队列卡死必须表现为失败。""" """限定等待时间,回归测试中的队列卡死必须表现为失败。"""
events = [] events = []
async with asyncio.timeout(10): while True:
while True: # 不用 asyncio.timeout部署环境是 Python 3.10wait_for 行为等价。
event = await ws.receive_json() event = await asyncio.wait_for(ws.receive_json(), timeout=10)
events.append(event) events.append(event)
if event["type"] == until: if event["type"] == until:
return events return events
def get_app(self) -> web.Application: def get_app(self) -> web.Application:
app = web.Application() app = web.Application()
@ -210,9 +216,10 @@ class WebSocketFlowTests(AioHTTPTestCase):
await ws.receive() await ws.receive()
else: else:
await ws.close() await ws.close()
async with asyncio.timeout(2): for _ in range(200):
while auxiliary.reset_speaker_session.await_count < 2: if auxiliary.reset_speaker_session.await_count >= 2:
await asyncio.sleep(0.01) break
await asyncio.sleep(0.01)
self.assertEqual(auxiliary.reset_speaker_session.await_count, 2) self.assertEqual(auxiliary.reset_speaker_session.await_count, 2)
async def test_extended_wav_header_is_removed_before_asr(self): async def test_extended_wav_header_is_removed_before_asr(self):
@ -346,3 +353,169 @@ class WebSocketFlowTests(AioHTTPTestCase):
self.assertEqual([item["speaker_id"] for item in latest], [0, 1, 0]) self.assertEqual([item["speaker_id"] for item in latest], [0, 1, 0])
self.assertTrue(all(item["speaker_evidence"] == "fresh" for item in latest)) self.assertTrue(all(item["speaker_evidence"] == "fresh" for item in latest))
await ws.close() await ws.close()
def _constant_pcm(ms: int, amplitude: int) -> bytes:
"""定幅 PCM16A 用高幅 6000、B 用低幅 1200两者都高于 VAD 门限。"""
return struct.pack("<h", amplitude) * (16 * ms)
def _silence_pcm(ms: int) -> bytes:
return b"\x00\x00" * (16 * ms)
def _speaker_frame_counts(audio: bytes) -> tuple[int, int]:
"""按帧 RMS 分类计数 (A, B)≥3000 视为说话人 A450~2999 视为 B。"""
counts = [0, 0]
usable = len(audio) - len(audio) % 640
for start in range(0, usable, 640):
frame = audio[start:start + 640]
squares = [int.from_bytes(frame[i:i + 2], "little", signed=True) ** 2 for i in range(0, 640, 2)]
rms = math.sqrt(sum(squares) / len(squares))
if rms >= 3000:
counts[0] += 1
elif rms >= 450:
counts[1] += 1
return counts[0], counts[1]
class _DummyWs:
closed = True
class SwitchModelService:
"""final 文本携带主导说话人和音频毫秒数,便于断言切段边界。"""
native_partial_supported = False
config = SimpleNamespace(base_url="http://fake/v1", model="fake-model")
async def transcribe(self, audio_bytes: bytes, source: str, file_name: str, partial: bool) -> str:
if partial:
return "partial text"
a, b = _speaker_frame_counts(audio_bytes)
if a + b == 0:
return ""
return f"{'A' if a >= b else 'B'} {round(len(audio_bytes) / 32)}ms"
class SwitchAuxiliaryService:
"""用帧能量成分构造二维 embeddingA→[1,0]、B→[0,1],混合按占比归一。"""
speaker_embedding_unsupported = False
def __init__(self, fail_windows: bool = False) -> None:
self.embedding_calls = 0
self.resolves: list[dict[str, object]] = []
self.fail_windows = fail_windows
async def speaker_embedding(self, pcm_bytes: bytes) -> list[float] | None:
self.embedding_calls += 1
if self.fail_windows:
return None
a, b = _speaker_frame_counts(pcm_bytes)
if a + b == 0:
return None
norm = math.hypot(a, b)
return [a / norm, b / norm]
async def resolve_speaker(
self,
audio_bytes: bytes,
session_id: str,
start_time_ms: float,
end_time_ms: float,
speaker_verified: bool = True,
) -> dict[str, object]:
_ = session_id
a, b = _speaker_frame_counts(audio_bytes)
speaker_id = 0 if a >= b else 1
self.resolves.append({
"audio_bytes": len(audio_bytes),
"speaker_verified": speaker_verified,
"start_time_ms": start_time_ms,
"end_time_ms": end_time_ms,
})
return {
"speaker_id": speaker_id,
"speaker_name": f"说话人 {speaker_id + 1}",
"speaker_evidence": "fresh",
"speaker_confidence": 0.9,
"speaker_strategy": "online_embedding_cluster_match",
"speaker_status": "confirmed",
}
async def reset_speaker_session(self, session_id: str) -> None:
_ = session_id
class SpeakerSwitchSegmentationTests(unittest.IsolatedAsyncioTestCase):
"""回归诊断出的核心缺陷A→B 停顿小于静音阈值时必须在窗口边界强制切段。"""
async def _run(self, pcm: bytes, auxiliary: SwitchAuxiliaryService) -> tuple[RealtimeSession, list]:
session = RealtimeSession(_DummyWs(), SwitchModelService(), auxiliary, {"source": "mic"})
session.audio_queue.put_nowait(pcm)
session.audio_queue.put_nowait(EOF)
await asyncio.wait_for(session.process_audio(), timeout=15)
jobs = []
while not session.speaker_queue.empty():
jobs.append(session.speaker_queue.get_nowait())
for job in jobs:
await session._resolve_speaker(job)
return session, jobs
@staticmethod
def _finals(session: RealtimeSession) -> list[dict]:
return [s for s in session.assembler.raw_snapshot() if s.get("sentence_type") == 1]
async def test_short_gap_handover_splits_at_speaker_change(self):
"""A 2.2s + 400ms 停顿(低于 800ms 阈值)+ B 2.2s:必须切成两个 final。"""
auxiliary = SwitchAuxiliaryService()
pcm = _constant_pcm(2200, 6000) + _silence_pcm(400) + _constant_pcm(2200, 1200) + _silence_pcm(1000)
session, jobs = await self._run(pcm, auxiliary)
finals = self._finals(session)
# 旧行为:整段合成一个 final 归 B新行为在 400ms 停顿后的起音处
# 2.6s切开A 的整句完整保留B 从自己的第一个音素开始。
self.assertEqual([s["sentence"] for s in finals], ["A 2600ms", "B 2200ms"])
self.assertEqual(finals[0]["commit_reason"], "speaker_change")
self.assertEqual(finals[0]["end_time"], 2600)
self.assertEqual(finals[1]["start_time"], 2600)
self.assertEqual([s["speaker_id"] for s in finals], [0, 1])
# 头段边界可能混入失配窗内容,排队标记为未验证;尾段正常。
self.assertEqual(
[(job.speaker_verified, len(job.audio)) for job in jobs],
[(False, 83200), (True, 70400)],
)
self.assertGreaterEqual(auxiliary.embedding_calls, 3)
async def test_same_speaker_continuation_is_not_split(self):
"""同一说话人 3s 后正常停顿:只出一个 final窗比对不误伤。"""
auxiliary = SwitchAuxiliaryService()
pcm = _constant_pcm(3000, 6000) + _silence_pcm(1000)
session, jobs = await self._run(pcm, auxiliary)
finals = self._finals(session)
self.assertEqual([s["sentence"] for s in finals], ["A 3000ms"])
self.assertEqual(finals[0]["commit_reason"], "silence")
self.assertTrue(all(job.speaker_verified for job in jobs))
self.assertGreaterEqual(auxiliary.embedding_calls, 2)
async def test_window_failures_disable_checks_without_blocking_asr(self):
"""窗提取连续失败 3 次后停用并告警一次ASR 照常产出 final。"""
auxiliary = SwitchAuxiliaryService(fail_windows=True)
pcm = _constant_pcm(2600, 6000) + _silence_pcm(1000)
session, jobs = await self._run(pcm, auxiliary)
self.assertFalse(session.window_check_available)
self.assertEqual(auxiliary.embedding_calls, 3)
self.assertTrue(session.speaker_warning_sent)
finals = self._finals(session)
self.assertEqual([s["sentence"] for s in finals], ["A 2600ms"])
self.assertEqual([job.speaker_verified for job in jobs], [True])
async def test_unsupported_endpoint_disables_after_single_call(self):
"""旧版辅助服务没有窗端点:探测一次即停用,不反复冲击服务。"""
auxiliary = SwitchAuxiliaryService(fail_windows=True)
auxiliary.speaker_embedding_unsupported = True
pcm = _constant_pcm(2600, 6000) + _silence_pcm(1000)
session, _ = await self._run(pcm, auxiliary)
self.assertFalse(session.window_check_available)
self.assertEqual(auxiliary.embedding_calls, 1)
self.assertEqual([s["sentence"] for s in self._finals(session)], ["A 2600ms"])

View File

@ -352,14 +352,29 @@ 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 speaker_embedding(self, audio_path: str) -> Any:
"""为一个短窗口提取归一化声纹,不读取也不更新任何会话聚类状态。
返回值与 _extract_embedding_sync 一致音频不足时为 None"""
async with self.inference_lock:
embedding = await asyncio.to_thread(self._extract_embedding_sync, audio_path)
if embedding is None:
return None
# 与 resolve 路径同样的防御:客户端阈值语义依赖单位向量。
return self._normalize_embedding(embedding)
async def resolve_speaker( async def resolve_speaker(
self, self,
audio_path: str, audio_path: str,
session_id: str, session_id: str,
start_time_ms: float, start_time_ms: float,
end_time_ms: float, end_time_ms: float,
speaker_verified: bool = True,
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""对一个实时 turn 提取声纹,并更新该 session 的在线聚类中心。""" """对一个实时 turn 提取声纹,并更新该 session 的在线聚类中心。
speaker_verified=False turn 来自换人强制切段的边界仍允许
匹配标签但不更新簇质心也不建立新簇避免混合音频污染聚类"""
async with self.inference_lock: async with self.inference_lock:
# 异常断网时客户端可能来不及 reset过期状态在下一次请求时回收。 # 异常断网时客户端可能来不及 reset过期状态在下一次请求时回收。
now = time.monotonic() now = time.monotonic()
@ -388,29 +403,45 @@ class AuxiliaryRuntime:
best_score = score best_score = score
best_cluster = cluster best_cluster = cluster
centroid_updated = False
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"]) if speaker_verified:
best_cluster["embedding"] = self._normalize_embedding( count = int(best_cluster["count"])
(best_cluster["embedding"] * count) + embedding best_cluster["embedding"] = self._normalize_embedding(
) (best_cluster["embedding"] * count) + embedding
best_cluster["count"] = count + 1 )
best_cluster["count"] = count + 1
centroid_updated = True
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 speaker_verified:
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})
centroid_updated = True
confidence = 0.75 confidence = 0.75
strategy = "online_embedding_cluster_new" strategy = "online_embedding_cluster_new"
else:
# 疑似混合音频不建立新簇:给临时编号并保持低于可信阈,
# 由展示层维持“未知”,直到该说话人的已验证片段来定簇。
speaker_id = len(clusters)
confidence = 0.55
strategy = "online_embedding_cluster_suspect_mixture"
return { return {
"speaker_id": speaker_id, "speaker_id": speaker_id,
"speaker_name": f"说话人 {speaker_id + 1}", "speaker_name": f"说话人 {speaker_id + 1}",
"speaker_evidence": "fresh", "speaker_evidence": "fresh",
"speaker_confidence": round(max(0.6, min(1.0, confidence)), 3), "speaker_confidence": round(max(0.5, min(1.0, confidence)), 3),
"speaker_strategy": strategy, "speaker_strategy": strategy,
"speaker_status": "confirmed", "speaker_status": "confirmed",
"speaker_reason": "当前片段独立声纹已完成在线聚类", "speaker_reason": (
"当前片段独立声纹已完成在线聚类"
if speaker_verified
else "换人边界片段仅匹配标签,未更新簇质心"
),
"speaker_verified": speaker_verified,
"speaker_centroid_updated": centroid_updated,
"start_time": start_time_ms, "start_time": start_time_ms,
"end_time": end_time_ms, "end_time": end_time_ms,
} }
@ -439,7 +470,8 @@ async def health_handler(request: web.Request) -> web.Response:
return web.json_response( return web.json_response(
{ {
"ready": ready, "ready": ready,
"speaker_protocol_version": 2, # v3resolve 支持 speaker_verified 质心卫生,并新增 /v1/speaker/embedding。
"speaker_protocol_version": 3,
"vad_model": vad_model_id, "vad_model": vad_model_id,
"vad_ready": vad_model_id is not None, "vad_ready": vad_model_id is not None,
"device": AUXILIARY_DEVICE, "device": AUXILIARY_DEVICE,
@ -550,6 +582,11 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
end_time_ms = _parse_form_float(form.get("end_time_ms"), "end_time_ms", default=start_time_ms) end_time_ms = _parse_form_float(form.get("end_time_ms"), "end_time_ms", default=start_time_ms)
except ValueError: except ValueError:
return web.json_response({"error": "turn time fields must be numbers"}, status=400) return web.json_response({"error": "turn time fields must be numbers"}, status=400)
# 旧客户端不带该字段时默认已验证,保持原语义;非法值同样回退为已验证。
verified_raw = form.get("speaker_verified", "1")
if isinstance(verified_raw, bytes):
verified_raw = verified_raw.decode("utf-8", "ignore")
speaker_verified = True if not isinstance(verified_raw, str) else verified_raw.strip().lower() not in {"0", "false", "no", "off"}
temp_path: str | None = None temp_path: str | None = None
try: try:
@ -561,6 +598,7 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
session_id, session_id,
start_time_ms, start_time_ms,
end_time_ms, end_time_ms,
speaker_verified,
) )
return web.json_response(result or { return web.json_response(result or {
"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0, "speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0,
@ -579,6 +617,32 @@ 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_embedding_handler(request: web.Request) -> web.Response:
"""接收活跃 turn 的短音频窗口,只返回归一化 embedding不触碰在线聚类状态。
WebSocket 侧用它做换人感知切段窗口向量绝不进入簇质心"""
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)
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
embedding = await runtime.speaker_embedding(temp_path)
if embedding is None:
return web.json_response({"ok": True, "embedding": None, "reason": "insufficient_audio"})
return web.json_response({"ok": True, "embedding": [float(value) for value in embedding.tolist()]})
except Exception as exc:
print(f"[speaker] window embedding failed: 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]
@ -599,6 +663,7 @@ async def create_app() -> web.Application:
app.router.add_post("/v1/vad", vad_handler) app.router.add_post("/v1/vad", vad_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/embedding", speaker_embedding_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

@ -165,5 +165,38 @@ class OnlineSpeakerTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(result["speaker_id"], -1) self.assertEqual(result["speaker_id"], -1)
self.assertEqual(runtime.speaker_clusters["test"][0]["count"], 1) self.assertEqual(runtime.speaker_clusters["test"][0]["count"], 1)
async def test_unverified_turn_labels_without_touching_clusters(self):
"""换人边界段:允许匹配既有簇,但不更新质心,也不建立新簇。"""
runtime = AuxiliaryRuntime()
vectors = iter(([1, 0], [0.98, 0.02], [0.2, 0.98], [0.2, 0.98]))
runtime._extract_embedding_sync = lambda _: np.array(next(vectors), dtype=np.float32)
await runtime.resolve_speaker("a.wav", "test", 0, 1000)
centroid_before = runtime.speaker_clusters["test"][0]["embedding"].copy()
matched = await runtime.resolve_speaker("head.wav", "test", 1000, 2000, speaker_verified=False)
self.assertEqual(matched["speaker_id"], 0)
self.assertTrue(matched["speaker_evidence"] == "fresh")
self.assertFalse(matched["speaker_centroid_updated"])
np.testing.assert_allclose(runtime.speaker_clusters["test"][0]["embedding"], centroid_before, atol=1e-6)
suspect = await runtime.resolve_speaker("mixed.wav", "test", 2000, 3000, speaker_verified=False)
self.assertLess(suspect["speaker_confidence"], 0.6)
self.assertEqual(suspect["speaker_strategy"], "online_embedding_cluster_suspect_mixture")
self.assertEqual(len(runtime.speaker_clusters["test"]), 1)
# 该说话人已验证的后续片段仍应能正常定簇。
verified = await runtime.resolve_speaker("b.wav", "test", 3000, 4000)
self.assertEqual(verified["speaker_id"], 1)
self.assertTrue(verified["speaker_centroid_updated"])
async def test_window_embedding_never_touches_cluster_state(self):
"""/v1/speaker/embedding 只出向量:单位化且零聚类副作用。"""
runtime = AuxiliaryRuntime()
runtime._extract_embedding_sync = lambda _: np.array([3, 4], dtype=np.float32)
embedding = await runtime.speaker_embedding("window.wav")
self.assertEqual(len(embedding), 2)
self.assertAlmostEqual(float(np.linalg.norm(embedding)), 1.0, places=6)
self.assertEqual(runtime.speaker_clusters, {})
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()