[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈
This commit is contained in:
+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