[独立唤醒与转写显示]:完成运行时本地唤醒分离,包含wake注入、转写输出和污染回归测试
This commit is contained in:
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user