[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈
This commit is contained in:
@@ -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:
|
||||
self.transcripts.append(text)
|
||||
self.events.append(f"transcript:{text}")
|
||||
if final:
|
||||
self.transcripts.append(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__":
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user