[独立唤醒与转写显示]:完成运行时本地唤醒分离,包含wake注入、转写输出和污染回归测试

This commit is contained in:
mkbk
2026-06-17 20:36:10 +08:00
parent e565164e6e
commit 4b21e0c346
3 changed files with 104 additions and 68 deletions
@@ -18,11 +18,11 @@
## 3. Runtime 独立唤醒与实时转写 ## 3. Runtime 独立唤醒与实时转写
- [ ] 3.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。 - [x] 3.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。
- [ ] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。 - [x] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。
- [ ] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。 - [x] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。
- [ ] 3.4 增加唤醒词污染回归测试;前置条件:3.2 完成;验收标准:LLM user message 不含 wake 音频文本;测试要点:fake wake metadata;优先级:P0;预计:30 分钟。 - [x] 3.4 增加唤醒词污染回归测试;前置条件:3.2 完成;验收标准:LLM user message 不含 wake 音频文本;测试要点:fake wake metadata;优先级:P0;预计:30 分钟。
- [ ] 3.5 验证并提交“独立唤醒与转写显示”模块;前置条件:3.1 至 3.4 完成;验收标准:compileall、unittest、security-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 - [x] 3.5 验证并提交“独立唤醒与转写显示”模块;前置条件:3.1 至 3.4 完成;验收标准:compileall、unittest、security-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。
## 4. 文档、真实验收与归档 ## 4. 文档、真实验收与归档
+51 -54
View File
@@ -8,17 +8,21 @@ from .config import AppConfig
from .conversation import ConversationContext from .conversation import ConversationContext
from .llm import OpenAICompatibleLlmProvider from .llm import OpenAICompatibleLlmProvider
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, VadRecorder from .vad import EnergyVadProvider, VadRecorder
from .wakeword import SherpaOnnxKeywordWakeWordProvider
class RuntimeReporter(Protocol): class RuntimeReporter(Protocol):
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
... ...
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
...
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
... ...
@@ -28,6 +32,11 @@ class TerminalRuntimeReporter:
prefix = f"[第{turn_id}轮] " if turn_id is not None else "" prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
print(f"{prefix}{message}", flush=True) print(f"{prefix}{message}", flush=True)
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 "转写中"
print(f"{prefix}{label}{text}", flush=True)
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
prefix = f"[第{turn_id}轮] " if turn_id is not None else "" prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True) print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True)
@@ -56,6 +65,7 @@ class LiveVoiceRuntime:
*, *,
config: AppConfig, config: AppConfig,
transport: AudioTransport, transport: AudioTransport,
wakeword: WakeWordProvider,
vad_recorder: VadRecorder, vad_recorder: VadRecorder,
stt: SttProvider, stt: SttProvider,
llm: LlmProvider, llm: LlmProvider,
@@ -66,6 +76,7 @@ class LiveVoiceRuntime:
) -> None: ) -> None:
self.config = config self.config = config
self.transport = transport self.transport = transport
self.wakeword = wakeword
self.vad_recorder = vad_recorder self.vad_recorder = vad_recorder
self.stt = stt self.stt = stt
self.llm = llm self.llm = llm
@@ -76,6 +87,7 @@ class LiveVoiceRuntime:
self._states: list[PipelineState] = [] self._states: list[PipelineState] = []
def load(self) -> None: def load(self) -> None:
self.wakeword.load()
self.vad_recorder.provider.load() self.vad_recorder.provider.load()
self.stt.load() self.stt.load()
self.tts.load() self.tts.load()
@@ -126,42 +138,38 @@ class LiveVoiceRuntime:
self.transport.stop() self.transport.stop()
def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError: def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError:
wake_error = self._wait_for_local_wake(turn_id)
if wake_error is not None:
return wake_error
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
if isinstance(user_segment, ProviderError):
return user_segment
self._state(PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id)
transcript = self.stt.transcribe(user_segment)
user_text = transcript.normalized_text
if not is_valid_transcript_text(user_text):
return ProviderError(
ErrorCode.STT_EMPTY_TRANSCRIPT,
"STT produced no meaningful user text",
True,
"live-runtime",
"stt",
)
self.reporter.transcript(user_text, final=True, turn_id=turn_id)
return user_text
def _wait_for_local_wake(self, turn_id: int) -> ProviderError | None:
self.wakeword.reset()
while True: while True:
wake_segment = self._capture_segment(turn_id, state_message="待机:检测到语音,正在判断唤醒词") frames = self.transport.read_frames(timeout_ms=100)
if isinstance(wake_segment, ProviderError): if not frames:
if wake_segment.code == ErrorCode.VAD_TIMEOUT_NO_SPEECH:
self._state(PipelineState.WAKE_LISTENING, "待机:继续等待唤醒词“小杰小杰”", turn_id=turn_id)
continue
return wake_segment
try:
wake_transcript = self.stt.transcribe(wake_segment).normalized_text
except ProviderError as exc:
if exc.code == ErrorCode.STT_EMPTY_TRANSCRIPT:
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未听清唤醒词", turn_id=turn_id)
continue
raise
remainder = _text_after_wake_word(wake_transcript, self.config.wake_word)
if remainder is None:
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未命中唤醒词", turn_id=turn_id)
continue continue
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id) for frame in frames:
if is_valid_transcript_text(remainder): event = self.wakeword.detect(frame)
return remainder if event is not None:
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题") self.wakeword.reset()
if isinstance(user_segment, ProviderError): return None
return user_segment
self._state(PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id)
transcript = self.stt.transcribe(user_segment)
user_text = transcript.normalized_text
if not is_valid_transcript_text(user_text):
return ProviderError(
ErrorCode.STT_EMPTY_TRANSCRIPT,
"STT produced no meaningful user text",
True,
"live-runtime",
"stt",
)
return user_text
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError: def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset() self.vad_recorder.reset()
@@ -180,7 +188,7 @@ class LiveVoiceRuntime:
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult: def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
self.context.append_user(user_text) self.context.append_user(user_text)
self._state(PipelineState.THINKING, f"思考中:{user_text}", turn_id=turn_id) self._state(PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
assistant_text = "" assistant_text = ""
try: try:
for delta in self.llm.stream_reply(self.context.build_llm_messages()): for delta in self.llm.stream_reply(self.context.build_llm_messages()):
@@ -237,6 +245,13 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
return LiveVoiceRuntime( return LiveVoiceRuntime(
config=config, config=config,
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device), transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
wakeword=SherpaOnnxKeywordWakeWordProvider(
config.speech_models_dir,
keyword=config.wake_word,
keywords_file=config.wake_keywords_file,
threshold=config.wake_kws_threshold,
score=config.wake_kws_score,
),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000),
stt=stt, stt=stt,
llm=OpenAICompatibleLlmProvider(config), llm=OpenAICompatibleLlmProvider(config),
@@ -247,21 +262,3 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
), ),
reporter=reporter, reporter=reporter,
) )
def _text_after_wake_word(text: str, wake_word: str) -> str | None:
compact_text = _compact(text)
compact_wake = _compact(wake_word)
index = compact_text.find(compact_wake)
if index < 0:
return None
end = index + len(compact_wake)
compact_remainder = compact_text[end:].strip(",。.!!?? ")
if not compact_remainder:
return ""
original = text.replace(" ", "")
return original[-len(compact_remainder) :]
def _compact(text: str) -> str:
return "".join(ch for ch in text.strip() if not ch.isspace())
+48 -9
View File
@@ -10,6 +10,7 @@ from owner_voice_pet.runtime import LiveVoiceRuntime
from owner_voice_pet.transport import MemoryAudioTransport from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SineTtsProvider from owner_voice_pet.tts import SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder from owner_voice_pet.vad import EnergyVadProvider, 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) -> list[AudioFrame]:
@@ -21,6 +22,17 @@ def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]:
] ]
def wake_frame(frame_id: int, timestamp_ms: int) -> AudioFrame:
return AudioFrame(
b"\xff\x7f",
16000,
1,
timestamp_ms,
frame_id,
{"duration_ms": 20, "wake_word": "小杰小杰", "wake_confidence": 0.95},
)
class QueueSttProvider: class QueueSttProvider:
def __init__(self, texts: list[str]) -> None: def __init__(self, texts: list[str]) -> None:
self.texts = list(texts) self.texts = list(texts)
@@ -39,10 +51,17 @@ class QueueSttProvider:
class RecordingReporter: class RecordingReporter:
def __init__(self) -> None: def __init__(self) -> None:
self.statuses: list[str] = [] self.statuses: list[str] = []
self.transcripts: list[str] = []
self.errors: list[str] = [] self.errors: list[str] = []
self.events: list[str] = []
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
self.statuses.append(message) self.statuses.append(message)
self.events.append(f"status:{message}")
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
self.transcripts.append(text)
self.events.append(f"transcript:{text}")
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
self.errors.append(f"{stage}:{code}:{message}") self.errors.append(f"{stage}:{code}:{message}")
@@ -50,8 +69,11 @@ class RecordingReporter:
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]: def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
frames = [] frames = []
for idx in range(4): for idx, _text in enumerate(texts):
frames.extend(segment_frames(idx * 4, idx * 80)) 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))
transport = MemoryAudioTransport(frames) transport = MemoryAudioTransport(frames)
stt = QueueSttProvider(texts) stt = QueueSttProvider(texts)
llm = MockLlmProvider(["这是答复。"]) llm = MockLlmProvider(["这是答复。"])
@@ -60,6 +82,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
runtime = LiveVoiceRuntime( runtime = LiveVoiceRuntime(
config=AppConfig(llm_api_key="secret", speech_provider="cloud"), config=AppConfig(llm_api_key="secret", speech_provider="cloud"),
transport=transport, transport=transport,
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=stt, stt=stt,
llm=llm, llm=llm,
@@ -72,19 +95,18 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
class LiveRuntimeTests(unittest.TestCase): class LiveRuntimeTests(unittest.TestCase):
def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None: def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None:
runtime, stt, llm, transport, reporter = make_runtime( runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"])
["小杰小杰", "第一问", "小杰小杰", "第二问"]
)
summary = runtime.run(max_turns=2) summary = runtime.run(max_turns=2)
self.assertEqual(summary.completed_turns, 2) self.assertEqual(summary.completed_turns, 2)
self.assertEqual(len(stt.calls), 4) self.assertEqual(len(stt.calls), 2)
self.assertEqual(len(llm.calls), 2) self.assertEqual(len(llm.calls), 2)
self.assertEqual(len(transport.played_segments), 2) self.assertEqual(len(transport.played_segments), 2)
self.assertEqual(reporter.transcripts, ["第一问", "第二问"])
self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1]) self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1])
def test_temporary_context_is_sent_to_second_llm_call(self) -> None: def test_temporary_context_is_sent_to_second_llm_call(self) -> None:
runtime, _, llm, _, _ = make_runtime(["小杰小杰", "第一问", "小杰小杰", "第二问"]) runtime, _, llm, _, _ = make_runtime(["第一问", "第二问"])
runtime.run(max_turns=2) runtime.run(max_turns=2)
second_call_text = [message.content for message in llm.calls[1]] second_call_text = [message.content for message in llm.calls[1]]
@@ -94,14 +116,31 @@ class LiveRuntimeTests(unittest.TestCase):
def test_new_runtime_context_starts_empty(self) -> None: def test_new_runtime_context_starts_empty(self) -> None:
first_context = ConversationContext() first_context = ConversationContext()
first_runtime, _, _, _, _ = make_runtime(["小杰小杰", "第一问"], context=first_context) first_runtime, _, _, _, _ = make_runtime(["第一问"], context=first_context)
first_runtime.run(max_turns=1) first_runtime.run(max_turns=1)
self.assertGreater(len(first_context.messages()), 0) self.assertGreater(len(first_context.messages()), 0)
second_context = ConversationContext() second_context = ConversationContext()
make_runtime(["小杰小杰", "第二问"], context=second_context) make_runtime(["第二问"], context=second_context)
self.assertEqual(second_context.messages(), ()) self.assertEqual(second_context.messages(), ())
def test_transcript_is_reported_before_llm_thinking(self) -> None:
runtime, _, _, _, reporter = make_runtime(["第一问"])
runtime.run(max_turns=1)
transcript_index = reporter.events.index("transcript:第一问")
thinking_index = next(
index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复"
)
self.assertLess(transcript_index, thinking_index)
def test_wake_keyword_does_not_pollute_llm_user_message(self) -> None:
runtime, _, llm, _, _ = make_runtime(["第一问"])
runtime.run(max_turns=1)
self.assertEqual(llm.calls[0][-1].content, "第一问")
self.assertNotIn("小杰小杰", llm.calls[0][-1].content)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()