[独立唤醒与转写显示]:完成运行时本地唤醒分离,包含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
+51 -54
View File
@@ -8,17 +8,21 @@ from .config import AppConfig
from .conversation import ConversationContext
from .llm import OpenAICompatibleLlmProvider
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 .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, VadRecorder
from .wakeword import SherpaOnnxKeywordWakeWordProvider
class RuntimeReporter(Protocol):
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:
...
@@ -28,6 +32,11 @@ class TerminalRuntimeReporter:
prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
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:
prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True)
@@ -56,6 +65,7 @@ class LiveVoiceRuntime:
*,
config: AppConfig,
transport: AudioTransport,
wakeword: WakeWordProvider,
vad_recorder: VadRecorder,
stt: SttProvider,
llm: LlmProvider,
@@ -66,6 +76,7 @@ class LiveVoiceRuntime:
) -> None:
self.config = config
self.transport = transport
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.llm = llm
@@ -76,6 +87,7 @@ class LiveVoiceRuntime:
self._states: list[PipelineState] = []
def load(self) -> None:
self.wakeword.load()
self.vad_recorder.provider.load()
self.stt.load()
self.tts.load()
@@ -126,42 +138,38 @@ class LiveVoiceRuntime:
self.transport.stop()
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:
wake_segment = self._capture_segment(turn_id, state_message="待机:检测到语音,正在判断唤醒词")
if isinstance(wake_segment, ProviderError):
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)
frames = self.transport.read_frames(timeout_ms=100)
if not frames:
continue
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
if is_valid_transcript_text(remainder):
return remainder
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",
)
return user_text
for frame in frames:
event = self.wakeword.detect(frame)
if event is not None:
self.wakeword.reset()
return None
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset()
@@ -180,7 +188,7 @@ class LiveVoiceRuntime:
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
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 = ""
try:
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(
config=config,
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),
stt=stt,
llm=OpenAICompatibleLlmProvider(config),
@@ -247,21 +262,3 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
),
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())