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