[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈

This commit is contained in:
mkbk
2026-06-17 22:16:56 +08:00
parent a64bb86da4
commit a77a172412
16 changed files with 331 additions and 17 deletions
+1
View File
@@ -2,6 +2,7 @@ OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1
OWNER_LLM_API_KEY=
OWNER_LLM_MODEL=mimo-v2.5
OWNER_LLM_API_STYLE=chat_completions
OWNER_REALTIME_TRANSCRIPT_ENABLED=1
OWNER_AUDIO_INPUT_DEVICE=
OWNER_AUDIO_OUTPUT_DEVICE=
OWNER_ASSET_DIR=assets/pet
+4 -1
View File
@@ -34,6 +34,7 @@ OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1
OWNER_LLM_API_KEY=
OWNER_LLM_MODEL=mimo-v2.5
OWNER_LLM_API_STYLE=chat_completions
OWNER_REALTIME_TRANSCRIPT_ENABLED=1
OWNER_WAKE_PROVIDER=local_kws
OWNER_WAKE_KEYWORDS_FILE=
OWNER_WAKE_KWS_THRESHOLD=0.15
@@ -62,6 +63,8 @@ OWNER_CONTEXT_MODE=session_memory
`OWNER_WAKE_PROVIDER=local_kws` 表示唤醒词“小杰小杰”由本地模型检测。唤醒命中后会先本地播报 `OWNER_WAKE_ACK_TEXT=我在`,再开始听取问题。`OWNER_SPEECH_PROVIDER=cloud` 只影响唤醒后的正式问题 ASR/TTS:它会把 VAD 切出来的用户问题片段发送到云端 ASR,不会上传连续麦克风流,也不会用云端判断唤醒。改成 `local` 时使用项目 `models/` 下的本地语音模型路径。
`OWNER_REALTIME_TRANSCRIPT_ENABLED=1` 表示录音期间会使用本地 streaming STT 实时显示中间转写,终端会输出 `实时转写:...`;最终发送给 LLM 的内容仍以 `转写结果:...` 为准。设置为 `0` 可以临时关闭实时显示。
## 本地模型
即使默认走云端 ASR/TTS,也必须准备本地唤醒模型;同一脚本会下载 wake、VAD、STT 模型,便于切换到 `OWNER_SPEECH_PROVIDER=local`
@@ -101,7 +104,7 @@ python3.11 scripts/download_speech_models.py --dir models
.venv/bin/python -m owner_voice_pet run-live
```
运行后终端状态来自 pipeline event bus,会显示待机、唤醒命中、应答中、请说出问题、录音中、检测到用户语音、用户语音结束、转写中、转写结果、思考中、播放中、恢复待机等状态。说“小杰小杰”,听到“我在”且看到“请说出问题”后再提问;听到回复后可以再次说“小杰小杰”继续下一轮。本次进程内会携带临时历史,程序退出后不保存。
运行后终端状态来自 pipeline event bus,会显示待机、唤醒命中、应答中、请说出问题、录音中、检测到用户语音、实时转写、用户语音结束、转写中、转写结果、思考中、播放中、恢复待机等状态。说“小杰小杰”,听到“我在”且看到“请说出问题”后再提问;听到回复后可以再次说“小杰小杰”继续下一轮。本次进程内会携带临时历史,程序退出后不保存。
## 验证
@@ -111,6 +111,16 @@ The first implementation is a per-turn heuristic endpoint, not persistent voicep
4. 主说话人画像就绪后,`OWNER_SPEAKER_ABSENT_MS` 是结束正式问题采集的主条件;主说话人连续缺席达到该值后直接进入 STT,不再额外等待普通 VAD 最小时长。
5. 普通 VAD 静音仍作为画像不足或音色特征不可用时的兜底,最大录音时长仍作为最终保护。
## Realtime Transcript Revision
录音期间新增 partial transcript 通道,用于解决“说话时看不到文字结果”的体验问题:
1. Pipeline 事件新增 `transcript_partial`。终端 reporter 将其显示为 `实时转写:<文本>`,最终结果仍显示为 `转写结果:<文本>`
2. partial transcript 只用于用户反馈,不写入 `ConversationContext`,不触发 LLMLLM 仍只消费 final STT 结果。
3. 当 final ASR/TTS 走 cloud 时,partial transcript 使用本地 `sherpa-onnx` streaming STT,避免对云端 ASR 高频请求。
4. `VoiceAssistantPipeline``speech_started` 后把 capture 阶段已开始录音的帧 feed 给 realtime STT session;当 partial 文本变化时才 emit,避免刷屏。
5. `OWNER_REALTIME_TRANSCRIPT_ENABLED=1` 默认启用;设置为 `0` 可临时回退到只显示 final transcript。
## Model Files
```text
@@ -141,6 +151,7 @@ models/
8. Pipeline stage failure: emit `stage_error`, recover to standby, and keep the process alive unless startup dependencies are missing.
9. Playback drain misconfiguration: negative `OWNER_POST_PLAYBACK_DRAIN_MS` remains invalid; non-zero values are treated as explicit user tuning rather than default behavior.
10. Speaker profile threshold misconfiguration: non-positive `OWNER_SPEAKER_PROFILE_MIN_MS` fails config validation.
11. Realtime STT failure: startup model缺失按 `model-check` 暴露;capture 中 partial session 失败不得把 partial 文本写入上下文。
## Testing Strategy
@@ -155,6 +166,8 @@ models/
9. First utterance preservation test validates that frames immediately after ACK are not discarded by post-playback drain.
10. Low-latency endpoint test validates that a short first question ends by primary speaker absence without waiting for a repeated second question.
11. Transport batching test validates that queued SoundDevice frames are returned together.
12. Partial transcript event test validates that realtime text appears after speech start and before final transcript.
13. Context isolation test validates partial transcript does not enter LLM messages.
## Migration
@@ -51,6 +51,7 @@
9. 首句保留问题:真人验收显示唤醒应答后仍会感觉“第一句话没有获取到”,当前 ACK 后会执行 `flush_input -> read/drop -> flush_input`,默认额外丢弃 50 ms 麦克风输入,用户若紧跟提示开口会损失正式问题开头。
10. 实时消费问题:`SoundDeviceAudioTransport.read_frames()` 每次只返回一个队列帧,在音频回调批量积压时会增加 pipeline 对真实麦克风流的追帧成本。
11. 画像门槛问题:`PrimarySpeakerVadRecorder._profile_ready()` 把主说话人画像就绪阈值绑定到 `OWNER_VAD_MIN_DURATION_MS`,默认至少等待 250 ms 后主说话人端点才参与结束判断,短句用户会被迫等普通 VAD 静音或重复说话。
12. 实时转写问题:当前终端只在整段录音结束并完成 final STT 后显示“转写结果”,用户说话期间看不到任何文字反馈,无法判断系统是否已经听到并识别当前句子。
## 详细需求
@@ -74,6 +75,9 @@
16. SoundDevice 音频输入 SHALL 支持一次读取当前队列中可用的多个帧,避免 pipeline 在真实麦克风输入积压时逐帧追赶。
17. 主说话人画像就绪 SHALL 使用独立配置 `OWNER_SPEAKER_PROFILE_MIN_MS`,默认 `120` ms;该阈值不得被 `OWNER_VAD_MIN_DURATION_MS` 放大。
18. 主说话人端点在画像就绪后 SHALL 以 `OWNER_SPEAKER_ABSENT_MS` 作为主要结束条件;主说话人连续缺席达到配置值后 SHALL 结束采集,不得额外等待普通 VAD 的最小时长门槛。
19. 录音期间 SHALL 支持 partial transcript 事件;当本地 streaming STT 产生新的中间文本时,终端 SHALL 立即显示 `实时转写:<文本>`
20. partial transcript SHALL 只作为用户可见反馈,不得直接写入对话上下文;LLM 输入仍以最终 `transcript_final` 文本为准。
21.`OWNER_SPEECH_PROVIDER=cloud` 时,partial transcript SHALL 使用本地 streaming STT,避免对云端 ASR 进行高频请求。
### 非功能需求
@@ -115,6 +119,7 @@
12. `OWNER_CONTEXT_MODE=session_memory`:本次进程内临时上下文。
13. `OWNER_POST_PLAYBACK_DRAIN_MS=0`:ACK 或 TTS 播放完成后只 flush 已积压输入,不额外读取并丢弃新音频。
14. `OWNER_SPEAKER_PROFILE_MIN_MS=120`:主说话人画像参与端点判断的最低有效语音长度。
15. `OWNER_REALTIME_TRANSCRIPT_ENABLED=1`:启用录音期间本地 streaming STT 中间结果显示。
终端输出:
@@ -222,6 +227,7 @@ standby
6. ACK 后首句保留:播放“我在”期间允许输入队列积压,播放完成后只执行一次队列 flush 清掉播放回声,不再额外读取 `post_playback_drain_ms` 毫秒并丢弃,默认值改为 0。
7. 批量读帧:真实 SoundDevice 输入在拿到首帧后立即 drain 当前队列中所有可用帧并返回给 pipeline,使 wake、capture 和 VAD 能在同一个循环内处理积压帧。
8. 快速主说话人结束:画像就绪最低语音长度由 `OWNER_SPEAKER_PROFILE_MIN_MS` 控制,默认 120 ms;一旦画像就绪,主说话人缺席计时达到 `OWNER_SPEAKER_ABSENT_MS` 即结束,不再叠加 `OWNER_VAD_MIN_DURATION_MS`
9. 实时转写显示:CaptureStage 在 `speech_started` 后把已录入的帧同时送入本地 streaming STT session;每当 partial 文本变化时发出 `transcript_partial`,终端显示 `实时转写:<文本>`;最终段落仍交给 configured STT provider 生成 `transcript_final`
### 数据库/状态管理变更
@@ -253,6 +259,8 @@ standby
| ACK 后额外丢弃音频截断首句 | 高 | 高 | 默认 `OWNER_POST_PLAYBACK_DRAIN_MS=0`;播放结束后只 flush 已积压输入;新增首句保留回归测试 |
| 主说话人画像等待过久 | 高 | 中 | 新增 `OWNER_SPEAKER_PROFILE_MIN_MS=120`;画像就绪后主说话人缺席结束不再等待普通 VAD 最小时长 |
| 麦克风帧队列积压导致状态滞后 | 中 | 中 | `SoundDeviceAudioTransport.read_frames()` 批量返回已积压帧;新增批量读帧测试 |
| partial STT 误识别影响 LLM | 中 | 中 | partial 只显示给用户,不进入 `ConversationContext`;最终 LLM 输入仍以 final STT 为准 |
| 本地 streaming STT 增加 CPU 占用 | 中 | 中 | 只在 capture started 且用户语音已开始后 feed;重复 partial 文本不重复输出;提供 `OWNER_REALTIME_TRANSCRIPT_ENABLED=0` 回退 |
| 手写音色特征不等于严格声纹识别 | 中 | 中 | 明确第一版为本轮临时主说话人端点;不承诺长期主人识别;保留后续接入 speaker embedding 模型的接口空间 |
| Pipeline 重构影响现有 CLI/测试 | 中 | 高 | 保留 `run-live` 命令和兼容类名;新增事件顺序测试、两轮回归测试和错误恢复测试 |
| 本地 KWS 增加启动加载时间 | 低 | 中 | 模型约 15 MB,启动加载一次;不在每轮重复加载 |
@@ -331,6 +339,7 @@ standby
7. `Live assistant pipeline events`:新增 stage 化事件要求,终端和后续 GUI 必须消费事件。
8. `Primary speaker endpointing`:新增本轮临时主说话人音色消失结束录音要求。
9. `Low latency capture and first utterance preservation`:新增 ACK 后不额外丢弃正式问题、批量读帧、独立画像就绪阈值和快速主说话人端点要求。
10. `Realtime partial transcript output`:新增录音期间 partial transcript 事件、终端显示和上下文隔离要求。
### 删除项
@@ -349,6 +358,7 @@ standby
5. M5:真人验收反馈修正唤醒提示顺序、KWS 阈值和 hybrid VAD,并在门禁通过后提交。
6. M6Stage 化 pipeline、事件总线、TurnController、主说话人端点和文档验收分模块提交。
7. M7:低延迟端点与首句保留修正完成后提交,保留真人 `run-live` 验收任务,不在用户确认前归档。
8. M8:录音期间实时转写显示完成后提交,保留真人 `run-live` 验收任务,不在用户确认前归档。
估时:
@@ -46,7 +46,15 @@ The live runtime SHALL preserve the start of the user's formal utterance after w
- **THEN** capture SHALL close the utterance and proceed to STT without waiting for a second repeated question or generic VAD minimum duration
### Requirement: Realtime transcript terminal output
The live runtime SHALL display the recognized user utterance text in the terminal after STT succeeds and before the LLM request is sent.
The live runtime SHALL display recognized user utterance text in the terminal both during recording when partial transcript is available and after final STT succeeds before the LLM request is sent.
#### Scenario: User utterance partial is available during recording
- **WHEN** the user is speaking after wake and the local realtime STT session produces a changed partial transcript
- **THEN** the terminal output SHALL include a realtime transcript message containing the partial text before final STT completes
#### Scenario: Partial transcript is displayed
- **WHEN** a partial transcript is emitted
- **THEN** it SHALL be treated as user-visible feedback only and SHALL NOT be appended to the conversation context or sent to the LLM
#### Scenario: User utterance is transcribed
- **WHEN** a live turn captures a user utterance and STT returns non-empty text
@@ -62,3 +62,11 @@
- [x] 8.3 实现主说话人画像独立就绪阈值;前置条件:8.2 完成;验收标准:新增 `OWNER_SPEAKER_PROFILE_MIN_MS=120`,主说话人缺席结束不再被 `OWNER_VAD_MIN_DURATION_MS` 拖慢;测试要点:短首句、背景噪声和重复提问不合并;优先级:P0;预计:45 分钟。
- [x] 8.4 补充配置文档、回归测试和门禁验证;前置条件:8.1 至 8.3 完成;验收标准:`.env.example`、README、配置显示和测试匹配新默认值;测试要点:compileall、unittest、security-check、model-check、device-check、OpenSpec strict;优先级:P0;预计:45 分钟。
- [x] 8.5 提交“低延迟端点”模块;前置条件:8.4 门禁通过;验收标准:中文 commit 信息为 `[低延迟端点]:完成首句保留和快速结束修正,包含ACK缓冲策略、批量读帧和主说话人端点回归测试`,提交后 `git status --short` 为空;优先级:P0;预计:10 分钟。
## 9. 录音期间实时转写显示
- [x] 9.1 更新 OpenSpec 描述 partial transcript;前置条件:用户要求“实时显示用户说话转换的结果”;验收标准:proposal/design/spec/tasks 明确录音期间输出 `transcript_partial`,最终 `transcript_final` 仍用于 LLM;测试要点:OpenSpec strict;优先级:P0;预计:30 分钟。
- [x] 9.2 实现 realtime transcript 事件和终端显示;前置条件:9.1 完成;验收标准:新增 `transcript_partial` 事件,终端输出 `实时转写:<文本>`,重复 partial 不刷屏;测试要点:事件顺序和 reporter 测试;优先级:P0;预计:45 分钟。
- [x] 9.3 实现本地 streaming STT partial provider;前置条件:本地 STT 模型已由 `model-check` 覆盖;验收标准:`SherpaOnnxSttProvider` 支持 streaming sessioncloud final ASR 模式下仍使用本地 streaming STT 做 partial;测试要点:fake session 与 metadata partial 单测;优先级:P0;预计:60 分钟。
- [x] 9.4 接入 `VoiceAssistantPipeline` 捕获循环;前置条件:9.2 至 9.3 完成;验收标准:用户说话期间持续 feed 已开始录音的帧,partial 在 `speech_started` 后、`transcript_final` 前输出;测试要点:live runtime partial 顺序测试;优先级:P0;预计:45 分钟。
- [x] 9.5 补充配置、README、门禁并提交;前置条件:9.1 至 9.4 完成;验收标准:新增 `OWNER_REALTIME_TRANSCRIPT_ENABLED=1`compileall、unittest、security-check、model-check、device-check、OpenSpec strict 通过后中文 commit;优先级:P0;预计:45 分钟。
+54 -2
View File
@@ -18,6 +18,7 @@ from .events import (
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TRANSCRIPT_PARTIAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
@@ -25,8 +26,16 @@ from .events import (
PipelineEventBus,
dispatch_pipeline_event,
)
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .models import AudioFrame, AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import (
AudioTransport,
LlmProvider,
RealtimeSttProvider,
RealtimeTranscriptSession,
SttProvider,
TtsProvider,
WakeWordProvider,
)
from .stt import is_valid_transcript_text
from .tts import SentenceBuffer
from .vad import VadRecorder
@@ -69,6 +78,7 @@ class TurnController:
wakeword: WakeWordProvider,
vad_recorder: VadRecorder,
stt: SttProvider,
realtime_stt: RealtimeSttProvider | None,
llm: LlmProvider,
tts: TtsProvider,
ack_tts: TtsProvider,
@@ -81,6 +91,7 @@ class TurnController:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts
@@ -146,6 +157,7 @@ class TurnController:
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset()
self.vad_recorder.provider.reset()
realtime_session = self._start_realtime_transcript()
self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id)
while True:
frames = self.transport.read_frames(timeout_ms=100)
@@ -158,7 +170,11 @@ class TurnController:
self._event(SPEECH_STARTED, PipelineState.RECORDING, "检测到用户语音", turn_id=turn_id)
if isinstance(result, ProviderError):
return result
if self.vad_recorder.started and realtime_session is not None:
self._emit_realtime_transcript(realtime_session, frame, turn_id)
if isinstance(result, AudioSegment):
if realtime_session is not None:
self._finish_realtime_transcript(realtime_session, turn_id)
self._event(
SPEECH_ENDED,
PipelineState.RECORDING,
@@ -168,6 +184,37 @@ class TurnController:
)
return result
def _start_realtime_transcript(self) -> RealtimeTranscriptSession | None:
if not self.config.realtime_transcript_enabled or self.realtime_stt is None:
return None
return self.realtime_stt.start_stream()
def _emit_realtime_transcript(
self,
realtime_session: RealtimeTranscriptSession,
frame: AudioFrame,
turn_id: int,
) -> None:
transcript = realtime_session.accept_frame(frame)
if transcript is None:
return
text = transcript.normalized_text
if not is_valid_transcript_text(text):
return
self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text})
def _finish_realtime_transcript(
self,
realtime_session: RealtimeTranscriptSession,
turn_id: int,
) -> None:
transcript = realtime_session.finish()
if transcript is None:
return
text = transcript.normalized_text
if is_valid_transcript_text(text):
self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text})
def _acknowledge_wake(self, turn_id: int) -> ProviderError | None:
text = self.config.wake_ack_text.strip()
if not text:
@@ -263,6 +310,7 @@ class VoiceAssistantPipeline:
llm: LlmProvider,
tts: TtsProvider,
context: ConversationContext,
realtime_stt: RealtimeSttProvider | None = None,
ack_tts: TtsProvider | None = None,
reporter: RuntimeReporter | None = None,
event_bus: PipelineEventBus | None = None,
@@ -273,6 +321,7 @@ class VoiceAssistantPipeline:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts or tts
@@ -288,6 +337,7 @@ class VoiceAssistantPipeline:
wakeword=wakeword,
vad_recorder=vad_recorder,
stt=stt,
realtime_stt=realtime_stt,
llm=llm,
tts=tts,
ack_tts=self.ack_tts,
@@ -300,6 +350,8 @@ class VoiceAssistantPipeline:
self.wakeword.load()
self.vad_recorder.provider.load()
self.stt.load()
if self.realtime_stt is not None and self.realtime_stt is not self.stt:
self.realtime_stt.load()
self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
+2
View File
@@ -51,6 +51,7 @@ def main(argv: list[str] | None = None) -> int:
"llm_model": config.llm_model,
"llm_api_style": config.llm_api_style,
"llm_stream": config.llm_stream,
"realtime_transcript_enabled": config.realtime_transcript_enabled,
"llm_api_key_present": bool(config.llm_api_key),
"asset_dir": str(config.asset_dir),
"wake_provider": config.wake_provider,
@@ -152,6 +153,7 @@ def main(argv: list[str] | None = None) -> int:
llm_model=config.llm_model,
llm_api_style=config.llm_api_style,
llm_stream=False,
realtime_transcript_enabled=config.realtime_transcript_enabled,
audio_input_device=config.audio_input_device,
audio_output_device=config.audio_output_device,
asset_dir=config.asset_dir,
+3
View File
@@ -16,6 +16,7 @@ class AppConfig:
llm_model: str = "mimo-v2.5"
llm_api_style: str = "chat_completions"
llm_stream: bool = True
realtime_transcript_enabled: bool = True
audio_input_device: str | None = None
audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet")
@@ -65,6 +66,8 @@ class AppConfig:
llm_model=get("LLM_MODEL", "mimo-v2.5") or "mimo-v2.5",
llm_api_style=get("LLM_API_STYLE", "chat_completions") or "chat_completions",
llm_stream=(get("LLM_STREAM", "1") or "1").lower() not in {"0", "false", "no"},
realtime_transcript_enabled=(get("REALTIME_TRANSCRIPT_ENABLED", "1") or "1").lower()
not in {"0", "false", "no"},
audio_input_device=get("AUDIO_INPUT_DEVICE"),
audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
+7 -2
View File
@@ -15,6 +15,7 @@ CAPTURE_STARTED = "capture_started"
SPEECH_STARTED = "speech_started"
SPEECH_ENDED = "speech_ended"
STT_STARTED = "stt_started"
TRANSCRIPT_PARTIAL = "transcript_partial"
TRANSCRIPT_FINAL = "transcript_final"
LLM_STARTED = "llm_started"
TTS_STARTED = "tts_started"
@@ -58,8 +59,12 @@ class PipelineEventBus:
def dispatch_pipeline_event(reporter: Any, event: PipelineEvent) -> None:
if event.type == TRANSCRIPT_FINAL:
reporter.transcript(str(event.payload.get("text", event.message)), final=True, turn_id=event.turn_id)
if event.type in {TRANSCRIPT_PARTIAL, TRANSCRIPT_FINAL}:
reporter.transcript(
str(event.payload.get("text", event.message)),
final=event.type == TRANSCRIPT_FINAL,
turn_id=event.turn_id,
)
return
if event.type == STAGE_ERROR:
error = event.payload.get("error")
+13
View File
@@ -68,6 +68,19 @@ class SttProvider(Protocol):
...
class RealtimeTranscriptSession(Protocol):
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
...
def finish(self) -> Transcript | None:
...
class RealtimeSttProvider(SttProvider, Protocol):
def start_stream(self) -> RealtimeTranscriptSession:
...
class LlmProvider(Protocol):
def stream_reply(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]:
...
+9 -2
View File
@@ -29,7 +29,7 @@ from .events import (
)
from .llm import OpenAICompatibleLlmProvider
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .protocols import AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider, WakeWordProvider
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
@@ -58,7 +58,7 @@ class TerminalRuntimeReporter:
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
label = "转写结果" if final else "转写"
label = "转写结果" if final else "实时转写"
print(f"{prefix}{label}{text}", flush=True)
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
@@ -336,6 +336,12 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider()
realtime_stt: RealtimeSttProvider | None = None
if config.realtime_transcript_enabled:
if isinstance(stt, SherpaOnnxSttProvider):
realtime_stt = stt
else:
realtime_stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
if config.vad_provider == "hybrid":
vad_provider = HybridVadProvider(
SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold),
@@ -375,6 +381,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
),
vad_recorder=recorder_cls(**recorder_kwargs),
stt=stt,
realtime_stt=realtime_stt,
llm=OpenAICompatibleLlmProvider(config),
tts=tts,
ack_tts=MacSayTtsProvider(),
+100 -1
View File
@@ -13,7 +13,7 @@ from pathlib import Path
from typing import Any
from .config import AppConfig
from .models import AudioSegment, ErrorCode, ProviderError, Transcript
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, Transcript
from .speech_models import stt_model_paths
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
@@ -58,6 +58,36 @@ class MetadataSttProvider:
raw_metadata=dict(segment.metadata),
)
def start_stream(self) -> "MetadataRealtimeTranscriptSession":
return MetadataRealtimeTranscriptSession(self.language)
class MetadataRealtimeTranscriptSession:
def __init__(self, language: str = "zh") -> None:
self.language = language
self._last_text = ""
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
text = str(
frame.metadata.get("partial_transcript")
or frame.metadata.get("transcript")
or ""
).strip()
if not is_valid_transcript_text(text) or text == self._last_text:
return None
self._last_text = text
return Transcript(
text=text,
language=str(frame.metadata.get("language", self.language)),
confidence=float(frame.metadata.get("stt_confidence", 1.0)),
duration_ms=int(frame.metadata.get("duration_ms", 20)),
provider="metadata-stt-partial",
raw_metadata=dict(frame.metadata),
)
def finish(self) -> Transcript | None:
return None
class CloudAsrSttProvider:
def __init__(
@@ -209,6 +239,17 @@ class SherpaOnnxSttProvider:
) from exc
self.loaded = True
def start_stream(self) -> "SherpaOnnxRealtimeTranscriptSession":
if not self.loaded or self._recognizer is None:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx STT provider is not loaded",
False,
"sherpa-onnx-stt",
"stt",
)
return SherpaOnnxRealtimeTranscriptSession(self._recognizer, self.language)
def transcribe(self, segment: AudioSegment) -> Transcript:
if not self.loaded or self._recognizer is None:
raise ProviderError(
@@ -264,6 +305,64 @@ def _segment_to_float32(segment: AudioSegment, np: Any) -> Any:
return samples
class SherpaOnnxRealtimeTranscriptSession:
def __init__(self, recognizer: Any, language: str = "zh") -> None:
self.recognizer = recognizer
self.language = language
self.stream = recognizer.create_stream()
self._last_text = ""
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
try:
import numpy as np
samples = _frame_to_float32(frame, np)
if samples.size == 0:
return None
self.stream.accept_waveform(frame.sample_rate, samples)
while self.recognizer.is_ready(self.stream):
self.recognizer.decode_stream(self.stream)
text = _recognizer_result_text(self.recognizer, self.stream)
except Exception as exc:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
f"sherpa-onnx realtime transcription failed: {exc}",
True,
"sherpa-onnx-stt",
"stt",
) from exc
if not is_valid_transcript_text(text) or text == self._last_text:
return None
self._last_text = text
return Transcript(
text=text,
language=self.language,
confidence=None,
duration_ms=int(frame.metadata.get("duration_ms", 20)),
provider="sherpa-onnx-stt-partial",
)
def finish(self) -> Transcript | None:
return None
def _frame_to_float32(frame: AudioFrame, np: Any) -> Any:
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
if frame.channels > 1 and samples.size:
samples = samples.reshape(-1, frame.channels).mean(axis=1)
return samples
def _recognizer_result_text(recognizer: Any, stream: Any) -> str:
if hasattr(recognizer, "get_result"):
result = recognizer.get_result(stream)
else:
result = recognizer.get_result_all(stream)
if isinstance(result, str):
return result.strip()
return str(getattr(result, "text", "")).strip()
def _segment_to_wav_bytes(segment: AudioSegment) -> bytes:
buffer = io.BytesIO()
with wave.open(buffer, "wb") as wav:
+40 -7
View File
@@ -14,6 +14,7 @@ from owner_voice_pet.events import (
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TRANSCRIPT_PARTIAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
@@ -23,16 +24,24 @@ from owner_voice_pet.llm import MockLlmProvider
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
from owner_voice_pet.runtime import build_live_runtime
from owner_voice_pet.stt import MetadataSttProvider
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder
from owner_voice_pet.wakeword import KeywordWakeWordProvider
def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]:
def segment_frames(start_id: int, start_ms: int, partials: list[str] | None = None) -> list[AudioFrame]:
partials = partials or []
first_metadata: dict[str, object] = {"duration_ms": 20, "speech": True}
second_metadata: dict[str, object] = {"duration_ms": 20, "speech": True}
if len(partials) >= 1:
first_metadata["partial_transcript"] = partials[0]
if len(partials) >= 2:
second_metadata["partial_transcript"] = partials[1]
return [
AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, {"duration_ms": 20, "speech": True}),
AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, {"duration_ms": 20, "speech": True}),
AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, first_metadata),
AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, second_metadata),
AudioFrame(b"\x00\x00", 16000, 1, start_ms + 40, start_id + 2, {"duration_ms": 20, "speech": False}),
AudioFrame(b"\x00\x00", 16000, 1, start_ms + 60, start_id + 3, {"duration_ms": 20, "speech": False}),
]
@@ -68,6 +77,7 @@ class RecordingReporter:
def __init__(self) -> None:
self.statuses: list[str] = []
self.transcripts: list[str] = []
self.partials: list[str] = []
self.errors: list[str] = []
self.events: list[str] = []
@@ -76,20 +86,29 @@ class RecordingReporter:
self.events.append(f"status:{message}")
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
if final:
self.transcripts.append(text)
self.events.append(f"transcript:{text}")
self.events.append(f"transcript:final:{text}")
else:
self.partials.append(text)
self.events.append(f"transcript:partial:{text}")
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
self.errors.append(f"{stage}:{code}:{message}")
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
def make_runtime(
texts: list[str],
context: ConversationContext | None = None,
partial_texts: list[list[str]] | None = None,
) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
frames = []
for idx, _text in enumerate(texts):
base_id = idx * 5
base_ms = idx * 120
frames.append(wake_frame(base_id, base_ms))
frames.extend(segment_frames(base_id + 1, base_ms + 20))
partials = partial_texts[idx] if partial_texts and idx < len(partial_texts) else None
frames.extend(segment_frames(base_id + 1, base_ms + 20, partials=partials))
transport = MemoryAudioTransport(frames)
stt = QueueSttProvider(texts)
llm = MockLlmProvider(["这是答复。"])
@@ -102,6 +121,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=stt,
realtime_stt=MetadataSttProvider() if partial_texts is not None else None,
llm=llm,
tts=tts,
context=context or ConversationContext(),
@@ -170,12 +190,24 @@ class LiveRuntimeTests(unittest.TestCase):
runtime, _, _, _, reporter = make_runtime(["第一问"])
runtime.run(max_turns=1)
transcript_index = reporter.events.index("transcript:第一问")
transcript_index = reporter.events.index("transcript:final:第一问")
thinking_index = next(
index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复"
)
self.assertLess(transcript_index, thinking_index)
def test_realtime_transcript_is_reported_while_capturing(self) -> None:
runtime, _, llm, _, reporter = make_runtime(["第一问"], partial_texts=[["第一", "第一问"]])
runtime.run(max_turns=1)
self.assertEqual(reporter.partials, ["第一", "第一问"])
self.assertEqual(reporter.transcripts, ["第一问"])
self.assertEqual(llm.calls[0][-1].content, "第一问")
event_types = [event.type for event in runtime.event_bus.events]
self.assertLess(event_types.index(SPEECH_STARTED), event_types.index(TRANSCRIPT_PARTIAL))
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(SPEECH_ENDED))
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(TRANSCRIPT_FINAL))
def test_wake_keyword_does_not_pollute_llm_user_message(self) -> None:
runtime, _, llm, _, _ = make_runtime(["第一问"])
runtime.run(max_turns=1)
@@ -186,6 +218,7 @@ class LiveRuntimeTests(unittest.TestCase):
def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None:
runtime = build_live_runtime(AppConfig(llm_api_key="secret"))
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
self.assertIsNotNone(runtime.realtime_stt)
if __name__ == "__main__":
+1
View File
@@ -69,6 +69,7 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
self.assertEqual(config.llm_api_key, "secret-value")
self.assertEqual(config.llm_model, "test-model")
self.assertTrue(config.realtime_transcript_enabled)
self.assertEqual(config.wake_provider, "local_kws")
self.assertEqual(config.wake_kws_threshold, 0.15)
self.assertEqual(config.wake_kws_score, 1.0)
+56
View File
@@ -3,10 +3,12 @@ from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
from owner_voice_pet.config import AppConfig
from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from owner_voice_pet.speech_models import stt_model_paths
from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder
from owner_voice_pet.wakeword import (
KeywordWakeWordProvider,
@@ -246,6 +248,60 @@ class WakeVadSttTests(unittest.TestCase):
provider.load()
self.assertEqual(raised.exception.code, ErrorCode.STT_MODEL_MISSING)
def test_sherpa_stt_streaming_session_emits_partial_text(self) -> None:
class FakeResult:
def __init__(self, text: str) -> None:
self.text = text
class FakeStream:
def __init__(self) -> None:
self.ready = False
self.text = ""
def accept_waveform(self, sample_rate, samples) -> None:
self.ready = True
self.text = "" if not self.text else "你好"
class FakeRecognizer:
def create_stream(self):
return FakeStream()
def is_ready(self, stream) -> bool:
return stream.ready
def decode_stream(self, stream) -> None:
stream.ready = False
def get_result(self, stream):
return FakeResult(stream.text)
class FakeOnlineRecognizer:
@staticmethod
def from_transducer(**kwargs):
return FakeRecognizer()
class FakeSherpa:
OnlineRecognizer = FakeOnlineRecognizer
with tempfile.TemporaryDirectory() as tmp:
paths = stt_model_paths(Path(tmp))
for name, path in paths.items():
if name == "model_dir":
continue
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text("fake", encoding="utf-8")
provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa)
provider.load()
session = provider.start_stream()
first = session.accept_frame(make_frame(1, 0, speech=True))
second = session.accept_frame(make_frame(2, 20, speech=True))
self.assertIsNotNone(first)
self.assertIsNotNone(second)
assert first is not None and second is not None
self.assertEqual(first.text, "")
self.assertEqual(second.text, "你好")
def test_cloud_asr_posts_audio_transcription_request(self) -> None:
class FakeResponse:
def __enter__(self):