[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈

This commit is contained in:
mkbk
2026-06-17 22:16:56 +08:00
parent a64bb86da4
commit a77a172412
16 changed files with 331 additions and 17 deletions
+54 -2
View File
@@ -18,6 +18,7 @@ from .events import (
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TRANSCRIPT_PARTIAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
@@ -25,8 +26,16 @@ from .events import (
PipelineEventBus,
dispatch_pipeline_event,
)
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .models import AudioFrame, AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import (
AudioTransport,
LlmProvider,
RealtimeSttProvider,
RealtimeTranscriptSession,
SttProvider,
TtsProvider,
WakeWordProvider,
)
from .stt import is_valid_transcript_text
from .tts import SentenceBuffer
from .vad import VadRecorder
@@ -69,6 +78,7 @@ class TurnController:
wakeword: WakeWordProvider,
vad_recorder: VadRecorder,
stt: SttProvider,
realtime_stt: RealtimeSttProvider | None,
llm: LlmProvider,
tts: TtsProvider,
ack_tts: TtsProvider,
@@ -81,6 +91,7 @@ class TurnController:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts
@@ -146,6 +157,7 @@ class TurnController:
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset()
self.vad_recorder.provider.reset()
realtime_session = self._start_realtime_transcript()
self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id)
while True:
frames = self.transport.read_frames(timeout_ms=100)
@@ -158,7 +170,11 @@ class TurnController:
self._event(SPEECH_STARTED, PipelineState.RECORDING, "检测到用户语音", turn_id=turn_id)
if isinstance(result, ProviderError):
return result
if self.vad_recorder.started and realtime_session is not None:
self._emit_realtime_transcript(realtime_session, frame, turn_id)
if isinstance(result, AudioSegment):
if realtime_session is not None:
self._finish_realtime_transcript(realtime_session, turn_id)
self._event(
SPEECH_ENDED,
PipelineState.RECORDING,
@@ -168,6 +184,37 @@ class TurnController:
)
return result
def _start_realtime_transcript(self) -> RealtimeTranscriptSession | None:
if not self.config.realtime_transcript_enabled or self.realtime_stt is None:
return None
return self.realtime_stt.start_stream()
def _emit_realtime_transcript(
self,
realtime_session: RealtimeTranscriptSession,
frame: AudioFrame,
turn_id: int,
) -> None:
transcript = realtime_session.accept_frame(frame)
if transcript is None:
return
text = transcript.normalized_text
if not is_valid_transcript_text(text):
return
self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text})
def _finish_realtime_transcript(
self,
realtime_session: RealtimeTranscriptSession,
turn_id: int,
) -> None:
transcript = realtime_session.finish()
if transcript is None:
return
text = transcript.normalized_text
if is_valid_transcript_text(text):
self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text})
def _acknowledge_wake(self, turn_id: int) -> ProviderError | None:
text = self.config.wake_ack_text.strip()
if not text:
@@ -263,6 +310,7 @@ class VoiceAssistantPipeline:
llm: LlmProvider,
tts: TtsProvider,
context: ConversationContext,
realtime_stt: RealtimeSttProvider | None = None,
ack_tts: TtsProvider | None = None,
reporter: RuntimeReporter | None = None,
event_bus: PipelineEventBus | None = None,
@@ -273,6 +321,7 @@ class VoiceAssistantPipeline:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts or tts
@@ -288,6 +337,7 @@ class VoiceAssistantPipeline:
wakeword=wakeword,
vad_recorder=vad_recorder,
stt=stt,
realtime_stt=realtime_stt,
llm=llm,
tts=tts,
ack_tts=self.ack_tts,
@@ -300,6 +350,8 @@ class VoiceAssistantPipeline:
self.wakeword.load()
self.vad_recorder.provider.load()
self.stt.load()
if self.realtime_stt is not None and self.realtime_stt is not self.stt:
self.realtime_stt.load()
self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
+2
View File
@@ -51,6 +51,7 @@ def main(argv: list[str] | None = None) -> int:
"llm_model": config.llm_model,
"llm_api_style": config.llm_api_style,
"llm_stream": config.llm_stream,
"realtime_transcript_enabled": config.realtime_transcript_enabled,
"llm_api_key_present": bool(config.llm_api_key),
"asset_dir": str(config.asset_dir),
"wake_provider": config.wake_provider,
@@ -152,6 +153,7 @@ def main(argv: list[str] | None = None) -> int:
llm_model=config.llm_model,
llm_api_style=config.llm_api_style,
llm_stream=False,
realtime_transcript_enabled=config.realtime_transcript_enabled,
audio_input_device=config.audio_input_device,
audio_output_device=config.audio_output_device,
asset_dir=config.asset_dir,
+3
View File
@@ -16,6 +16,7 @@ class AppConfig:
llm_model: str = "mimo-v2.5"
llm_api_style: str = "chat_completions"
llm_stream: bool = True
realtime_transcript_enabled: bool = True
audio_input_device: str | None = None
audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet")
@@ -65,6 +66,8 @@ class AppConfig:
llm_model=get("LLM_MODEL", "mimo-v2.5") or "mimo-v2.5",
llm_api_style=get("LLM_API_STYLE", "chat_completions") or "chat_completions",
llm_stream=(get("LLM_STREAM", "1") or "1").lower() not in {"0", "false", "no"},
realtime_transcript_enabled=(get("REALTIME_TRANSCRIPT_ENABLED", "1") or "1").lower()
not in {"0", "false", "no"},
audio_input_device=get("AUDIO_INPUT_DEVICE"),
audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
+7 -2
View File
@@ -15,6 +15,7 @@ CAPTURE_STARTED = "capture_started"
SPEECH_STARTED = "speech_started"
SPEECH_ENDED = "speech_ended"
STT_STARTED = "stt_started"
TRANSCRIPT_PARTIAL = "transcript_partial"
TRANSCRIPT_FINAL = "transcript_final"
LLM_STARTED = "llm_started"
TTS_STARTED = "tts_started"
@@ -58,8 +59,12 @@ class PipelineEventBus:
def dispatch_pipeline_event(reporter: Any, event: PipelineEvent) -> None:
if event.type == TRANSCRIPT_FINAL:
reporter.transcript(str(event.payload.get("text", event.message)), final=True, turn_id=event.turn_id)
if event.type in {TRANSCRIPT_PARTIAL, TRANSCRIPT_FINAL}:
reporter.transcript(
str(event.payload.get("text", event.message)),
final=event.type == TRANSCRIPT_FINAL,
turn_id=event.turn_id,
)
return
if event.type == STAGE_ERROR:
error = event.payload.get("error")
+13
View File
@@ -68,6 +68,19 @@ class SttProvider(Protocol):
...
class RealtimeTranscriptSession(Protocol):
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
...
def finish(self) -> Transcript | None:
...
class RealtimeSttProvider(SttProvider, Protocol):
def start_stream(self) -> RealtimeTranscriptSession:
...
class LlmProvider(Protocol):
def stream_reply(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]:
...
+9 -2
View File
@@ -29,7 +29,7 @@ from .events import (
)
from .llm import OpenAICompatibleLlmProvider
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .protocols import AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider, WakeWordProvider
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
@@ -58,7 +58,7 @@ class TerminalRuntimeReporter:
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 "转写"
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:
@@ -336,6 +336,12 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider()
realtime_stt: RealtimeSttProvider | None = None
if config.realtime_transcript_enabled:
if isinstance(stt, SherpaOnnxSttProvider):
realtime_stt = stt
else:
realtime_stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
if config.vad_provider == "hybrid":
vad_provider = HybridVadProvider(
SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold),
@@ -375,6 +381,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
),
vad_recorder=recorder_cls(**recorder_kwargs),
stt=stt,
realtime_stt=realtime_stt,
llm=OpenAICompatibleLlmProvider(config),
tts=tts,
ack_tts=MacSayTtsProvider(),
+100 -1
View File
@@ -13,7 +13,7 @@ from pathlib import Path
from typing import Any
from .config import AppConfig
from .models import AudioSegment, ErrorCode, ProviderError, Transcript
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, Transcript
from .speech_models import stt_model_paths
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
@@ -58,6 +58,36 @@ class MetadataSttProvider:
raw_metadata=dict(segment.metadata),
)
def start_stream(self) -> "MetadataRealtimeTranscriptSession":
return MetadataRealtimeTranscriptSession(self.language)
class MetadataRealtimeTranscriptSession:
def __init__(self, language: str = "zh") -> None:
self.language = language
self._last_text = ""
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
text = str(
frame.metadata.get("partial_transcript")
or frame.metadata.get("transcript")
or ""
).strip()
if not is_valid_transcript_text(text) or text == self._last_text:
return None
self._last_text = text
return Transcript(
text=text,
language=str(frame.metadata.get("language", self.language)),
confidence=float(frame.metadata.get("stt_confidence", 1.0)),
duration_ms=int(frame.metadata.get("duration_ms", 20)),
provider="metadata-stt-partial",
raw_metadata=dict(frame.metadata),
)
def finish(self) -> Transcript | None:
return None
class CloudAsrSttProvider:
def __init__(
@@ -209,6 +239,17 @@ class SherpaOnnxSttProvider:
) from exc
self.loaded = True
def start_stream(self) -> "SherpaOnnxRealtimeTranscriptSession":
if not self.loaded or self._recognizer is None:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx STT provider is not loaded",
False,
"sherpa-onnx-stt",
"stt",
)
return SherpaOnnxRealtimeTranscriptSession(self._recognizer, self.language)
def transcribe(self, segment: AudioSegment) -> Transcript:
if not self.loaded or self._recognizer is None:
raise ProviderError(
@@ -264,6 +305,64 @@ def _segment_to_float32(segment: AudioSegment, np: Any) -> Any:
return samples
class SherpaOnnxRealtimeTranscriptSession:
def __init__(self, recognizer: Any, language: str = "zh") -> None:
self.recognizer = recognizer
self.language = language
self.stream = recognizer.create_stream()
self._last_text = ""
def accept_frame(self, frame: AudioFrame) -> Transcript | None:
try:
import numpy as np
samples = _frame_to_float32(frame, np)
if samples.size == 0:
return None
self.stream.accept_waveform(frame.sample_rate, samples)
while self.recognizer.is_ready(self.stream):
self.recognizer.decode_stream(self.stream)
text = _recognizer_result_text(self.recognizer, self.stream)
except Exception as exc:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
f"sherpa-onnx realtime transcription failed: {exc}",
True,
"sherpa-onnx-stt",
"stt",
) from exc
if not is_valid_transcript_text(text) or text == self._last_text:
return None
self._last_text = text
return Transcript(
text=text,
language=self.language,
confidence=None,
duration_ms=int(frame.metadata.get("duration_ms", 20)),
provider="sherpa-onnx-stt-partial",
)
def finish(self) -> Transcript | None:
return None
def _frame_to_float32(frame: AudioFrame, np: Any) -> Any:
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
if frame.channels > 1 and samples.size:
samples = samples.reshape(-1, frame.channels).mean(axis=1)
return samples
def _recognizer_result_text(recognizer: Any, stream: Any) -> str:
if hasattr(recognizer, "get_result"):
result = recognizer.get_result(stream)
else:
result = recognizer.get_result_all(stream)
if isinstance(result, str):
return result.strip()
return str(getattr(result, "text", "")).strip()
def _segment_to_wav_bytes(segment: AudioSegment) -> bytes:
buffer = io.BytesIO()
with wave.open(buffer, "wb") as wav: