[本地语音降噪]:完成本地语音链路和噪音过滤,包含降噪模型、本地ASR和实时字幕稳定策略

This commit is contained in:
mkbk
2026-06-17 22:55:33 +08:00
parent a77a172412
commit 8c75fc5baf
22 changed files with 803 additions and 64 deletions
+3
View File
@@ -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",
+13
View File
@@ -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()
+156
View File
@@ -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,
},
)
]
+8
View File
@@ -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,
+18 -2
View File
@@ -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(
+2
View File
@@ -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"
+14
View File
@@ -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:
...
+7
View File
@@ -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),
+48 -12
View File
@@ -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
View File
@@ -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(