[独立唤醒与转写显示]:完成运行时本地唤醒分离,包含wake注入、转写输出和污染回归测试
This commit is contained in:
@@ -18,11 +18,11 @@
|
||||
|
||||
## 3. Runtime 独立唤醒与实时转写
|
||||
|
||||
- [ ] 3.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。
|
||||
- [ ] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。
|
||||
- [ ] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。
|
||||
- [ ] 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.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。
|
||||
- [x] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。
|
||||
- [x] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。
|
||||
- [x] 3.4 增加唤醒词污染回归测试;前置条件:3.2 完成;验收标准:LLM user message 不含 wake 音频文本;测试要点:fake wake metadata;优先级:P0;预计:30 分钟。
|
||||
- [x] 3.5 验证并提交“独立唤醒与转写显示”模块;前置条件:3.1 至 3.4 完成;验收标准:compileall、unittest、security-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。
|
||||
|
||||
## 4. 文档、真实验收与归档
|
||||
|
||||
|
||||
@@ -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,27 +138,10 @@ class LiveVoiceRuntime:
|
||||
self.transport.stop()
|
||||
|
||||
def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError:
|
||||
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)
|
||||
continue
|
||||
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)
|
||||
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
|
||||
@@ -161,8 +156,21 @@ class LiveVoiceRuntime:
|
||||
"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:
|
||||
frames = self.transport.read_frames(timeout_ms=100)
|
||||
if not frames:
|
||||
continue
|
||||
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()
|
||||
self.vad_recorder.provider.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())
|
||||
|
||||
@@ -10,6 +10,7 @@ from owner_voice_pet.runtime import LiveVoiceRuntime
|
||||
from owner_voice_pet.transport import MemoryAudioTransport
|
||||
from owner_voice_pet.tts import SineTtsProvider
|
||||
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]:
|
||||
@@ -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:
|
||||
def __init__(self, texts: list[str]) -> None:
|
||||
self.texts = list(texts)
|
||||
@@ -39,10 +51,17 @@ class QueueSttProvider:
|
||||
class RecordingReporter:
|
||||
def __init__(self) -> None:
|
||||
self.statuses: list[str] = []
|
||||
self.transcripts: list[str] = []
|
||||
self.errors: list[str] = []
|
||||
self.events: list[str] = []
|
||||
|
||||
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
|
||||
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:
|
||||
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]:
|
||||
frames = []
|
||||
for idx in range(4):
|
||||
frames.extend(segment_frames(idx * 4, idx * 80))
|
||||
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))
|
||||
transport = MemoryAudioTransport(frames)
|
||||
stt = QueueSttProvider(texts)
|
||||
llm = MockLlmProvider(["这是答复。"])
|
||||
@@ -60,6 +82,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
|
||||
runtime = LiveVoiceRuntime(
|
||||
config=AppConfig(llm_api_key="secret", speech_provider="cloud"),
|
||||
transport=transport,
|
||||
wakeword=KeywordWakeWordProvider(),
|
||||
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
||||
stt=stt,
|
||||
llm=llm,
|
||||
@@ -72,19 +95,18 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
|
||||
|
||||
class LiveRuntimeTests(unittest.TestCase):
|
||||
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)
|
||||
|
||||
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(transport.played_segments), 2)
|
||||
self.assertEqual(reporter.transcripts, ["第一问", "第二问"])
|
||||
self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1])
|
||||
|
||||
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)
|
||||
|
||||
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:
|
||||
first_context = ConversationContext()
|
||||
first_runtime, _, _, _, _ = make_runtime(["小杰小杰", "第一问"], context=first_context)
|
||||
first_runtime, _, _, _, _ = make_runtime(["第一问"], context=first_context)
|
||||
first_runtime.run(max_turns=1)
|
||||
self.assertGreater(len(first_context.messages()), 0)
|
||||
|
||||
second_context = ConversationContext()
|
||||
make_runtime(["小杰小杰", "第二问"], context=second_context)
|
||||
make_runtime(["第二问"], context=second_context)
|
||||
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__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user