[本地语音降噪]:完成本地语音链路和噪音过滤,包含降噪模型、本地ASR和实时字幕稳定策略
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
|
||||
from .config import AppConfig
|
||||
from .assistant_pipeline import TurnController, VoiceAssistantPipeline
|
||||
from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor
|
||||
from .events import PipelineEvent, PipelineEventBus
|
||||
from .models import (
|
||||
AudioFrame,
|
||||
@@ -33,6 +34,8 @@ __all__ = [
|
||||
"AppConfig",
|
||||
"TurnController",
|
||||
"VoiceAssistantPipeline",
|
||||
"NoopAudioPreprocessor",
|
||||
"SherpaOnnxDenoiserPreprocessor",
|
||||
"PipelineEvent",
|
||||
"PipelineEventBus",
|
||||
"AudioFrame",
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from .audio_preprocess import NoopAudioPreprocessor
|
||||
from .config import AppConfig
|
||||
from .conversation import ConversationContext
|
||||
from .events import (
|
||||
@@ -28,6 +29,7 @@ from .events import (
|
||||
)
|
||||
from .models import AudioFrame, AudioSegment, ErrorCode, PipelineState, ProviderError
|
||||
from .protocols import (
|
||||
AudioPreprocessor,
|
||||
AudioTransport,
|
||||
LlmProvider,
|
||||
RealtimeSttProvider,
|
||||
@@ -77,6 +79,7 @@ class TurnController:
|
||||
transport: AudioTransport,
|
||||
wakeword: WakeWordProvider,
|
||||
vad_recorder: VadRecorder,
|
||||
audio_preprocessor: AudioPreprocessor,
|
||||
stt: SttProvider,
|
||||
realtime_stt: RealtimeSttProvider | None,
|
||||
llm: LlmProvider,
|
||||
@@ -90,6 +93,7 @@ class TurnController:
|
||||
self.transport = transport
|
||||
self.wakeword = wakeword
|
||||
self.vad_recorder = vad_recorder
|
||||
self.audio_preprocessor = audio_preprocessor
|
||||
self.stt = stt
|
||||
self.realtime_stt = realtime_stt
|
||||
self.llm = llm
|
||||
@@ -157,6 +161,7 @@ class TurnController:
|
||||
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
|
||||
self.vad_recorder.reset()
|
||||
self.vad_recorder.provider.reset()
|
||||
self.audio_preprocessor.reset()
|
||||
realtime_session = self._start_realtime_transcript()
|
||||
self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id)
|
||||
while True:
|
||||
@@ -164,6 +169,10 @@ class TurnController:
|
||||
if not frames:
|
||||
continue
|
||||
for frame in frames:
|
||||
try:
|
||||
frame = self.audio_preprocessor.process_frame(frame)
|
||||
except ProviderError as exc:
|
||||
return exc
|
||||
was_started = self.vad_recorder.started
|
||||
result = self.vad_recorder.feed(frame)
|
||||
if not was_started and self.vad_recorder.started:
|
||||
@@ -310,6 +319,7 @@ class VoiceAssistantPipeline:
|
||||
llm: LlmProvider,
|
||||
tts: TtsProvider,
|
||||
context: ConversationContext,
|
||||
audio_preprocessor: AudioPreprocessor | None = None,
|
||||
realtime_stt: RealtimeSttProvider | None = None,
|
||||
ack_tts: TtsProvider | None = None,
|
||||
reporter: RuntimeReporter | None = None,
|
||||
@@ -320,6 +330,7 @@ class VoiceAssistantPipeline:
|
||||
self.transport = transport
|
||||
self.wakeword = wakeword
|
||||
self.vad_recorder = vad_recorder
|
||||
self.audio_preprocessor = audio_preprocessor or NoopAudioPreprocessor()
|
||||
self.stt = stt
|
||||
self.realtime_stt = realtime_stt
|
||||
self.llm = llm
|
||||
@@ -336,6 +347,7 @@ class VoiceAssistantPipeline:
|
||||
transport=transport,
|
||||
wakeword=wakeword,
|
||||
vad_recorder=vad_recorder,
|
||||
audio_preprocessor=self.audio_preprocessor,
|
||||
stt=stt,
|
||||
realtime_stt=realtime_stt,
|
||||
llm=llm,
|
||||
@@ -349,6 +361,7 @@ class VoiceAssistantPipeline:
|
||||
def load(self) -> None:
|
||||
self.wakeword.load()
|
||||
self.vad_recorder.provider.load()
|
||||
self.audio_preprocessor.load()
|
||||
self.stt.load()
|
||||
if self.realtime_stt is not None and self.realtime_stt is not self.stt:
|
||||
self.realtime_stt.load()
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .models import AudioFrame, ErrorCode, ProviderError
|
||||
from .speech_models import denoiser_model_path
|
||||
|
||||
|
||||
class NoopAudioPreprocessor:
|
||||
def load(self) -> None:
|
||||
return None
|
||||
|
||||
def reset(self) -> None:
|
||||
return None
|
||||
|
||||
def process_frame(self, frame: AudioFrame) -> AudioFrame:
|
||||
return frame
|
||||
|
||||
def flush(self) -> list[AudioFrame]:
|
||||
return []
|
||||
|
||||
|
||||
class SherpaOnnxDenoiserPreprocessor:
|
||||
def __init__(
|
||||
self,
|
||||
models_dir: str | Path,
|
||||
*,
|
||||
sherpa_module: Any | None = None,
|
||||
) -> None:
|
||||
self.models_dir = Path(models_dir)
|
||||
self._sherpa = sherpa_module
|
||||
self._denoiser: Any | None = None
|
||||
self.loaded = False
|
||||
|
||||
def load(self) -> None:
|
||||
model = denoiser_model_path(self.models_dir)
|
||||
if not model.exists():
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_MODEL_MISSING,
|
||||
f"sherpa-onnx denoiser model is missing: {model}",
|
||||
False,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
)
|
||||
sherpa_onnx = self._sherpa
|
||||
if sherpa_onnx is None:
|
||||
try:
|
||||
import sherpa_onnx # type: ignore[import-not-found]
|
||||
except Exception as exc:
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_FAILED,
|
||||
f"sherpa_onnx is not available: {exc}",
|
||||
False,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
) from exc
|
||||
try:
|
||||
gtcrn = sherpa_onnx.OfflineSpeechDenoiserGtcrnModelConfig(model=str(model))
|
||||
model_config = sherpa_onnx.OfflineSpeechDenoiserModelConfig(
|
||||
gtcrn=gtcrn,
|
||||
num_threads=1,
|
||||
provider="cpu",
|
||||
)
|
||||
config = sherpa_onnx.OnlineSpeechDenoiserConfig(model=model_config)
|
||||
self._denoiser = sherpa_onnx.OnlineSpeechDenoiser(config)
|
||||
except Exception as exc:
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_FAILED,
|
||||
f"failed to load sherpa-onnx denoiser: {exc}",
|
||||
False,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
) from exc
|
||||
self.loaded = True
|
||||
|
||||
def reset(self) -> None:
|
||||
if self._denoiser is not None and hasattr(self._denoiser, "reset"):
|
||||
self._denoiser.reset()
|
||||
|
||||
def process_frame(self, frame: AudioFrame) -> AudioFrame:
|
||||
if not self.loaded or self._denoiser is None:
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_FAILED,
|
||||
"sherpa-onnx denoiser is not loaded",
|
||||
False,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
)
|
||||
try:
|
||||
import numpy as np
|
||||
|
||||
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)
|
||||
denoised = self._denoiser.run(samples, frame.sample_rate)
|
||||
output_samples = np.asarray(getattr(denoised, "samples"), dtype=np.float32)
|
||||
output_sample_rate = int(getattr(denoised, "sample_rate", frame.sample_rate))
|
||||
clipped = np.clip(output_samples, -1.0, 1.0)
|
||||
pcm = (clipped * 32767.0).astype(np.int16).tobytes()
|
||||
except ProviderError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_FAILED,
|
||||
f"sherpa-onnx denoiser failed: {exc}",
|
||||
True,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
) from exc
|
||||
metadata = dict(frame.metadata)
|
||||
metadata["denoised"] = True
|
||||
metadata["noise_filter_provider"] = "sherpa_onnx_gtcrn"
|
||||
return AudioFrame(
|
||||
pcm=pcm,
|
||||
sample_rate=output_sample_rate,
|
||||
channels=1,
|
||||
timestamp_ms=frame.timestamp_ms,
|
||||
frame_id=frame.frame_id,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
def flush(self) -> list[AudioFrame]:
|
||||
if self._denoiser is None or not hasattr(self._denoiser, "flush"):
|
||||
return []
|
||||
try:
|
||||
import numpy as np
|
||||
|
||||
denoised = self._denoiser.flush()
|
||||
samples = np.asarray(getattr(denoised, "samples"), dtype=np.float32)
|
||||
if samples.size == 0:
|
||||
return []
|
||||
sample_rate = int(getattr(denoised, "sample_rate", 16000))
|
||||
pcm = (np.clip(samples, -1.0, 1.0) * 32767.0).astype(np.int16).tobytes()
|
||||
except Exception as exc:
|
||||
raise ProviderError(
|
||||
ErrorCode.NOISE_FILTER_FAILED,
|
||||
f"sherpa-onnx denoiser flush failed: {exc}",
|
||||
True,
|
||||
"sherpa-onnx-gtcrn",
|
||||
"audio-preprocess",
|
||||
) from exc
|
||||
return [
|
||||
AudioFrame(
|
||||
pcm=pcm,
|
||||
sample_rate=sample_rate,
|
||||
channels=1,
|
||||
timestamp_ms=0,
|
||||
frame_id=0,
|
||||
metadata={
|
||||
"denoised": True,
|
||||
"noise_filter_provider": "sherpa_onnx_gtcrn",
|
||||
"flush": True,
|
||||
},
|
||||
)
|
||||
]
|
||||
@@ -6,6 +6,7 @@ import re
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
from .audio_preprocess import SherpaOnnxDenoiserPreprocessor
|
||||
from .assets import validate_pet_assets
|
||||
from .config import AppConfig
|
||||
from .conversation import ConversationContext
|
||||
@@ -62,6 +63,9 @@ def main(argv: list[str] | None = None) -> int:
|
||||
"post_playback_drain_ms": config.post_playback_drain_ms,
|
||||
"pipeline_mode": config.pipeline_mode,
|
||||
"endpoint_mode": config.endpoint_mode,
|
||||
"noise_filter_enabled": config.noise_filter_enabled,
|
||||
"noise_filter_provider": config.noise_filter_provider,
|
||||
"wake_denoise_enabled": config.wake_denoise_enabled,
|
||||
"speaker_profile_ms": config.speaker_profile_ms,
|
||||
"speaker_profile_min_ms": config.speaker_profile_min_ms,
|
||||
"speaker_absent_ms": config.speaker_absent_ms,
|
||||
@@ -113,6 +117,7 @@ def main(argv: list[str] | None = None) -> int:
|
||||
).load()
|
||||
SherpaOnnxVadProvider(models_dir).load()
|
||||
SherpaOnnxSttProvider(str(models_dir)).load()
|
||||
SherpaOnnxDenoiserPreprocessor(models_dir).load()
|
||||
provider_load_checked = True
|
||||
except ProviderError as exc:
|
||||
errors.append(exc)
|
||||
@@ -166,6 +171,9 @@ def main(argv: list[str] | None = None) -> int:
|
||||
post_playback_drain_ms=config.post_playback_drain_ms,
|
||||
pipeline_mode=config.pipeline_mode,
|
||||
endpoint_mode=config.endpoint_mode,
|
||||
noise_filter_enabled=config.noise_filter_enabled,
|
||||
noise_filter_provider=config.noise_filter_provider,
|
||||
wake_denoise_enabled=config.wake_denoise_enabled,
|
||||
speaker_profile_ms=config.speaker_profile_ms,
|
||||
speaker_profile_min_ms=config.speaker_profile_min_ms,
|
||||
speaker_absent_ms=config.speaker_absent_ms,
|
||||
|
||||
@@ -29,6 +29,9 @@ class AppConfig:
|
||||
post_playback_drain_ms: int = 0
|
||||
pipeline_mode: str = "live_turn_based"
|
||||
endpoint_mode: str = "primary_speaker"
|
||||
noise_filter_enabled: bool = True
|
||||
noise_filter_provider: str = "sherpa_onnx_gtcrn"
|
||||
wake_denoise_enabled: bool = False
|
||||
speaker_profile_ms: int = 600
|
||||
speaker_profile_min_ms: int = 120
|
||||
speaker_absent_ms: int = 300
|
||||
@@ -40,7 +43,7 @@ class AppConfig:
|
||||
vad_end_silence_ms: int = 350
|
||||
vad_no_speech_timeout_ms: int = 5000
|
||||
vad_max_recording_ms: int = 12000
|
||||
speech_provider: str = "cloud"
|
||||
speech_provider: str = "local"
|
||||
asr_model: str = "mimo-v2.5-asr"
|
||||
tts_model: str = "mimo-v2.5-tts"
|
||||
tts_voice: str = "mimo_default"
|
||||
@@ -80,6 +83,9 @@ class AppConfig:
|
||||
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "0") or "0"),
|
||||
pipeline_mode=(get("PIPELINE_MODE", "live_turn_based") or "live_turn_based").lower(),
|
||||
endpoint_mode=(get("ENDPOINT_MODE", "primary_speaker") or "primary_speaker").lower(),
|
||||
noise_filter_enabled=(get("NOISE_FILTER_ENABLED", "1") or "1").lower() not in {"0", "false", "no"},
|
||||
noise_filter_provider=(get("NOISE_FILTER_PROVIDER", "sherpa_onnx_gtcrn") or "sherpa_onnx_gtcrn").lower(),
|
||||
wake_denoise_enabled=(get("WAKE_DENOISE_ENABLED", "0") or "0").lower() in {"1", "true", "yes"},
|
||||
speaker_profile_ms=int(get("SPEAKER_PROFILE_MS", "600") or "600"),
|
||||
speaker_profile_min_ms=int(get("SPEAKER_PROFILE_MIN_MS", "120") or "120"),
|
||||
speaker_absent_ms=int(get("SPEAKER_ABSENT_MS", "300") or "300"),
|
||||
@@ -93,7 +99,7 @@ class AppConfig:
|
||||
vad_end_silence_ms=int(get("VAD_END_SILENCE_MS", "350") or "350"),
|
||||
vad_no_speech_timeout_ms=int(get("VAD_NO_SPEECH_TIMEOUT_MS", "5000") or "5000"),
|
||||
vad_max_recording_ms=int(get("VAD_MAX_RECORDING_MS", "12000") or "12000"),
|
||||
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
|
||||
speech_provider=(get("SPEECH_PROVIDER", "local") or "local").lower(),
|
||||
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr",
|
||||
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
|
||||
tts_voice=get("TTS_VOICE", "mimo_default") or "mimo_default",
|
||||
@@ -189,6 +195,16 @@ class AppConfig:
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.noise_filter_provider not in {"sherpa_onnx_gtcrn"}:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_NOISE_FILTER_PROVIDER must be sherpa_onnx_gtcrn",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.post_playback_drain_ms < 0:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
|
||||
@@ -41,6 +41,8 @@ class ErrorCode(str, Enum):
|
||||
TTS_MODEL_MISSING = "TTS_MODEL_MISSING"
|
||||
TTS_SYNTHESIS_FAILED = "TTS_SYNTHESIS_FAILED"
|
||||
TTS_EMPTY_AUDIO = "TTS_EMPTY_AUDIO"
|
||||
NOISE_FILTER_MODEL_MISSING = "NOISE_FILTER_MODEL_MISSING"
|
||||
NOISE_FILTER_FAILED = "NOISE_FILTER_FAILED"
|
||||
ASSET_MISSING = "ASSET_MISSING"
|
||||
VALIDATION_FAILED = "VALIDATION_FAILED"
|
||||
|
||||
|
||||
@@ -49,6 +49,20 @@ class WakeWordProvider(Protocol):
|
||||
...
|
||||
|
||||
|
||||
class AudioPreprocessor(Protocol):
|
||||
def load(self) -> None:
|
||||
...
|
||||
|
||||
def reset(self) -> None:
|
||||
...
|
||||
|
||||
def process_frame(self, frame: AudioFrame) -> AudioFrame:
|
||||
...
|
||||
|
||||
def flush(self) -> list[AudioFrame]:
|
||||
...
|
||||
|
||||
|
||||
class VadProvider(Protocol):
|
||||
def load(self) -> None:
|
||||
...
|
||||
|
||||
@@ -4,6 +4,7 @@ import sys
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor
|
||||
from .config import AppConfig
|
||||
from .assistant_pipeline import VoiceAssistantPipeline
|
||||
from .conversation import ConversationContext
|
||||
@@ -369,6 +370,11 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
|
||||
"min_rms": config.speaker_min_rms,
|
||||
}
|
||||
)
|
||||
audio_preprocessor = (
|
||||
SherpaOnnxDenoiserPreprocessor(config.speech_models_dir)
|
||||
if config.noise_filter_enabled
|
||||
else NoopAudioPreprocessor()
|
||||
)
|
||||
return VoiceAssistantPipeline(
|
||||
config=config,
|
||||
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
|
||||
@@ -380,6 +386,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
|
||||
score=config.wake_kws_score,
|
||||
),
|
||||
vad_recorder=recorder_cls(**recorder_kwargs),
|
||||
audio_preprocessor=audio_preprocessor,
|
||||
stt=stt,
|
||||
realtime_stt=realtime_stt,
|
||||
llm=OpenAICompatibleLlmProvider(config),
|
||||
|
||||
@@ -11,15 +11,20 @@ from .models import ErrorCode, ProviderError
|
||||
DEFAULT_VAD_URL = "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/silero_vad.onnx"
|
||||
DEFAULT_STT_URL = (
|
||||
"https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/"
|
||||
"sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23.tar.bz2"
|
||||
"sherpa-onnx-streaming-zipformer-ctc-zh-int8-2025-06-30.tar.bz2"
|
||||
)
|
||||
DEFAULT_STT_DIR = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23"
|
||||
DEFAULT_STT_DIR = "sherpa-onnx-streaming-zipformer-ctc-zh-int8-2025-06-30"
|
||||
DEFAULT_KWS_URL = (
|
||||
"https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/"
|
||||
"sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile.tar.bz2"
|
||||
)
|
||||
DEFAULT_KWS_DIR = "sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile"
|
||||
DEFAULT_KWS_KEYWORDS = "x iǎo j ié x iǎo j ié @小杰小杰\n"
|
||||
DEFAULT_DENOISER_URL = (
|
||||
"https://github.com/k2-fsa/sherpa-onnx/releases/download/"
|
||||
"speech-enhancement-models/gtcrn_simple.onnx"
|
||||
)
|
||||
DEFAULT_DENOISER_PATH = "denoise/gtcrn_simple.onnx"
|
||||
|
||||
REQUIRED_MODEL_FILES = (
|
||||
f"wake/{DEFAULT_KWS_DIR}/tokens.txt",
|
||||
@@ -29,9 +34,8 @@ REQUIRED_MODEL_FILES = (
|
||||
"wake/keywords.txt",
|
||||
"vad/silero_vad.onnx",
|
||||
f"stt/{DEFAULT_STT_DIR}/tokens.txt",
|
||||
f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx",
|
||||
f"stt/{DEFAULT_STT_DIR}/decoder-epoch-99-avg-1.onnx",
|
||||
f"stt/{DEFAULT_STT_DIR}/joiner-epoch-99-avg-1.int8.onnx",
|
||||
f"stt/{DEFAULT_STT_DIR}/model.int8.onnx",
|
||||
DEFAULT_DENOISER_PATH,
|
||||
)
|
||||
|
||||
|
||||
@@ -68,6 +72,7 @@ def default_manifest() -> dict[str, Any]:
|
||||
"wake": DEFAULT_KWS_URL,
|
||||
"vad": DEFAULT_VAD_URL,
|
||||
"stt": DEFAULT_STT_URL,
|
||||
"denoiser": DEFAULT_DENOISER_URL,
|
||||
},
|
||||
"providers": {
|
||||
"wake": {
|
||||
@@ -84,12 +89,14 @@ def default_manifest() -> dict[str, Any]:
|
||||
"path": "vad/silero_vad.onnx",
|
||||
},
|
||||
"stt": {
|
||||
"type": "sherpa-onnx-streaming-transducer",
|
||||
"type": "sherpa-onnx-streaming-zipformer2-ctc",
|
||||
"model_dir": f"stt/{DEFAULT_STT_DIR}",
|
||||
"tokens": f"stt/{DEFAULT_STT_DIR}/tokens.txt",
|
||||
"encoder": f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx",
|
||||
"decoder": f"stt/{DEFAULT_STT_DIR}/decoder-epoch-99-avg-1.onnx",
|
||||
"joiner": f"stt/{DEFAULT_STT_DIR}/joiner-epoch-99-avg-1.int8.onnx",
|
||||
"model": f"stt/{DEFAULT_STT_DIR}/model.int8.onnx",
|
||||
},
|
||||
"denoiser": {
|
||||
"type": "sherpa-onnx-gtcrn",
|
||||
"path": DEFAULT_DENOISER_PATH,
|
||||
},
|
||||
},
|
||||
"required_files": list(REQUIRED_MODEL_FILES),
|
||||
@@ -127,6 +134,13 @@ def vad_model_path(models_dir: str | Path) -> Path:
|
||||
return root / str(path)
|
||||
|
||||
|
||||
def denoiser_model_path(models_dir: str | Path) -> Path:
|
||||
root = Path(models_dir)
|
||||
manifest = load_manifest(root)
|
||||
path = manifest.get("providers", {}).get("denoiser", {}).get("path", DEFAULT_DENOISER_PATH)
|
||||
return root / str(path)
|
||||
|
||||
|
||||
def wake_model_paths(models_dir: str | Path) -> dict[str, Path]:
|
||||
root = Path(models_dir)
|
||||
manifest = load_manifest(root)
|
||||
@@ -144,12 +158,21 @@ def wake_model_paths(models_dir: str | Path) -> dict[str, Path]:
|
||||
}
|
||||
|
||||
|
||||
def stt_model_paths(model_path: str | Path) -> dict[str, Path]:
|
||||
def stt_model_paths(model_path: str | Path) -> dict[str, Any]:
|
||||
root = Path(model_path)
|
||||
if (root / "manifest.json").exists() or (root / "stt").exists():
|
||||
manifest = load_manifest(root)
|
||||
stt = manifest.get("providers", {}).get("stt", {})
|
||||
stt_type = str(stt.get("type", "sherpa-onnx-streaming-transducer"))
|
||||
if stt_type == "sherpa-onnx-streaming-zipformer2-ctc":
|
||||
return {
|
||||
"type": stt_type,
|
||||
"model_dir": root / str(stt.get("model_dir", f"stt/{DEFAULT_STT_DIR}")),
|
||||
"tokens": root / str(stt.get("tokens", f"stt/{DEFAULT_STT_DIR}/tokens.txt")),
|
||||
"model": root / str(stt.get("model", f"stt/{DEFAULT_STT_DIR}/model.int8.onnx")),
|
||||
}
|
||||
return {
|
||||
"type": stt_type,
|
||||
"model_dir": root / str(stt.get("model_dir", f"stt/{DEFAULT_STT_DIR}")),
|
||||
"tokens": root / str(stt.get("tokens", f"stt/{DEFAULT_STT_DIR}/tokens.txt")),
|
||||
"encoder": root / str(stt.get("encoder", f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx")),
|
||||
@@ -157,6 +180,7 @@ def stt_model_paths(model_path: str | Path) -> dict[str, Path]:
|
||||
"joiner": root / str(stt.get("joiner", f"stt/{DEFAULT_STT_DIR}/joiner-epoch-99-avg-1.int8.onnx")),
|
||||
}
|
||||
return {
|
||||
"type": "sherpa-onnx-streaming-transducer",
|
||||
"model_dir": root,
|
||||
"tokens": root / "tokens.txt",
|
||||
"encoder": root / "encoder-epoch-99-avg-1.int8.onnx",
|
||||
@@ -195,11 +219,23 @@ def model_status_errors(status: SpeechModelStatus) -> list[ProviderError]:
|
||||
"model-check",
|
||||
)
|
||||
)
|
||||
if status.missing_files:
|
||||
denoise_missing = tuple(item for item in status.missing_files if item.startswith("denoise/"))
|
||||
speech_missing = tuple(item for item in status.missing_files if not item.startswith("denoise/"))
|
||||
if speech_missing:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.STT_MODEL_MISSING,
|
||||
"missing speech model files: " + ", ".join(status.missing_files),
|
||||
"missing speech model files: " + ", ".join(speech_missing),
|
||||
False,
|
||||
"speech-models",
|
||||
"model-check",
|
||||
)
|
||||
)
|
||||
if denoise_missing:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.NOISE_FILTER_MODEL_MISSING,
|
||||
"missing denoiser model files: " + ", ".join(denoise_missing),
|
||||
False,
|
||||
"speech-models",
|
||||
"model-check",
|
||||
|
||||
+46
-13
@@ -17,12 +17,28 @@ from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, Transcri
|
||||
from .speech_models import stt_model_paths
|
||||
|
||||
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
|
||||
_PARTIAL_MIN_MEANINGFUL_CHARS = 3
|
||||
|
||||
|
||||
def is_valid_transcript_text(text: str) -> bool:
|
||||
return bool(_MEANINGFUL_TEXT.search(text.strip()))
|
||||
|
||||
|
||||
def _meaningful_text_length(text: str) -> int:
|
||||
return len(_MEANINGFUL_TEXT.findall(text.strip()))
|
||||
|
||||
|
||||
def should_emit_partial_transcript(text: str, last_text: str) -> bool:
|
||||
normalized = text.strip()
|
||||
if not is_valid_transcript_text(normalized) or _meaningful_text_length(normalized) < _PARTIAL_MIN_MEANINGFUL_CHARS:
|
||||
return False
|
||||
if normalized == last_text:
|
||||
return False
|
||||
if not last_text:
|
||||
return True
|
||||
return normalized.startswith(last_text)
|
||||
|
||||
|
||||
class MetadataSttProvider:
|
||||
def __init__(self, language: str = "zh") -> None:
|
||||
self.language = language
|
||||
@@ -73,7 +89,7 @@ class MetadataRealtimeTranscriptSession:
|
||||
or frame.metadata.get("transcript")
|
||||
or ""
|
||||
).strip()
|
||||
if not is_valid_transcript_text(text) or text == self._last_text:
|
||||
if not should_emit_partial_transcript(text, self._last_text):
|
||||
return None
|
||||
self._last_text = text
|
||||
return Transcript(
|
||||
@@ -197,7 +213,11 @@ class SherpaOnnxSttProvider:
|
||||
"stt",
|
||||
)
|
||||
paths = stt_model_paths(self.model_path)
|
||||
missing = [name for name, path in paths.items() if name != "model_dir" and not path.exists()]
|
||||
missing = [
|
||||
name
|
||||
for name, path in paths.items()
|
||||
if name not in {"type", "model_dir"} and isinstance(path, Path) and not path.exists()
|
||||
]
|
||||
if missing:
|
||||
raise ProviderError(
|
||||
ErrorCode.STT_MODEL_MISSING,
|
||||
@@ -219,16 +239,29 @@ class SherpaOnnxSttProvider:
|
||||
"stt",
|
||||
) from exc
|
||||
try:
|
||||
self._recognizer = sherpa_onnx.OnlineRecognizer.from_transducer(
|
||||
tokens=str(paths["tokens"]),
|
||||
encoder=str(paths["encoder"]),
|
||||
decoder=str(paths["decoder"]),
|
||||
joiner=str(paths["joiner"]),
|
||||
num_threads=1,
|
||||
decoding_method="greedy_search",
|
||||
enable_endpoint_detection=True,
|
||||
provider="cpu",
|
||||
)
|
||||
model_type = str(paths.get("type", "sherpa-onnx-streaming-transducer"))
|
||||
if model_type == "sherpa-onnx-streaming-zipformer2-ctc":
|
||||
self._recognizer = sherpa_onnx.OnlineRecognizer.from_zipformer2_ctc(
|
||||
tokens=str(paths["tokens"]),
|
||||
model=str(paths["model"]),
|
||||
num_threads=1,
|
||||
sample_rate=16000,
|
||||
feature_dim=80,
|
||||
enable_endpoint_detection=True,
|
||||
decoding_method="greedy_search",
|
||||
provider="cpu",
|
||||
)
|
||||
else:
|
||||
self._recognizer = sherpa_onnx.OnlineRecognizer.from_transducer(
|
||||
tokens=str(paths["tokens"]),
|
||||
encoder=str(paths["encoder"]),
|
||||
decoder=str(paths["decoder"]),
|
||||
joiner=str(paths["joiner"]),
|
||||
num_threads=1,
|
||||
decoding_method="greedy_search",
|
||||
enable_endpoint_detection=True,
|
||||
provider="cpu",
|
||||
)
|
||||
except Exception as exc:
|
||||
raise ProviderError(
|
||||
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||
@@ -331,7 +364,7 @@ class SherpaOnnxRealtimeTranscriptSession:
|
||||
"sherpa-onnx-stt",
|
||||
"stt",
|
||||
) from exc
|
||||
if not is_valid_transcript_text(text) or text == self._last_text:
|
||||
if not should_emit_partial_transcript(text, self._last_text):
|
||||
return None
|
||||
self._last_text = text
|
||||
return Transcript(
|
||||
|
||||
Reference in New Issue
Block a user