[唤醒灵敏度与端点恢复]:完成唤醒提示顺序和Hybrid VAD优化,包含阈值默认值、缓冲时序和测试覆盖

This commit is contained in:
mkbk
2026-06-17 21:00:28 +08:00
parent 44f708a3dc
commit 95e4996434
13 changed files with 145 additions and 29 deletions
+2 -1
View File
@@ -17,7 +17,7 @@ from .models import (
)
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
from .vad import EnergyVadProvider, VadRecorder
from .vad import EnergyVadProvider, HybridVadProvider, VadRecorder
from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .conversation import ConversationContext
from .llm import MockLlmProvider, OpenAICompatibleLlmProvider
@@ -38,6 +38,7 @@ __all__ = [
"KeywordWakeWordProvider",
"SherpaOnnxKeywordWakeWordProvider",
"EnergyVadProvider",
"HybridVadProvider",
"VadRecorder",
"CloudAsrSttProvider",
"MetadataSttProvider",
+8 -8
View File
@@ -22,11 +22,11 @@ class AppConfig:
log_dir: Path = Path("logs")
wake_provider: str = "local_kws"
wake_keywords_file: Path | None = None
wake_kws_threshold: float = 0.25
wake_kws_threshold: float = 0.15
wake_kws_score: float = 1.0
wake_ack_text: str = "我在"
post_playback_drain_ms: int = 250
vad_provider: str = "local"
post_playback_drain_ms: int = 50
vad_provider: str = "hybrid"
vad_threshold: float = 0.5
vad_min_duration_ms: int = 250
vad_end_silence_ms: int = 350
@@ -63,11 +63,11 @@ class AppConfig:
log_dir=Path(get("LOG_DIR", "logs") or "logs"),
wake_provider=(get("WAKE_PROVIDER", "local_kws") or "local_kws").lower(),
wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None,
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"),
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.15") or "0.15"),
wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"),
wake_ack_text=get("WAKE_ACK_TEXT", "我在") or "我在",
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "250") or "250"),
vad_provider=(get("VAD_PROVIDER", "local") or "local").lower(),
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "50") or "50"),
vad_provider=(get("VAD_PROVIDER", "hybrid") or "hybrid").lower(),
vad_threshold=float(get("VAD_THRESHOLD", "0.5") or "0.5"),
vad_min_duration_ms=int(get("VAD_MIN_DURATION_MS", "250") or "250"),
vad_end_silence_ms=int(get("VAD_END_SILENCE_MS", "350") or "350"),
@@ -178,11 +178,11 @@ class AppConfig:
"startup",
)
)
if self.vad_provider not in {"local", "energy"}:
if self.vad_provider not in {"hybrid", "local", "energy"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_PROVIDER must be local or energy",
"OWNER_VAD_PROVIDER must be hybrid, local, or energy",
False,
"config",
"startup",
+9 -3
View File
@@ -12,7 +12,7 @@ from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, Wa
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .vad import EnergyVadProvider, HybridVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import SherpaOnnxKeywordWakeWordProvider
@@ -145,10 +145,11 @@ class LiveVoiceRuntime:
wake_error = self._wait_for_local_wake(turn_id)
if wake_error is not None:
return wake_error
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中", turn_id=turn_id)
ack_error = self._acknowledge_wake(turn_id)
if ack_error is not None:
return ack_error
self._state(PipelineState.SPEECH_DETECTING, "请说出问题", turn_id=turn_id)
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
if isinstance(user_segment, ProviderError):
return user_segment
@@ -277,7 +278,12 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider()
if config.vad_provider == "local":
if config.vad_provider == "hybrid":
vad_provider = HybridVadProvider(
SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold),
EnergyVadProvider(),
)
elif config.vad_provider == "local":
vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold)
else:
vad_provider = EnergyVadProvider()
+47
View File
@@ -157,6 +157,53 @@ class SherpaOnnxVadProvider:
self._model.reset()
class HybridVadProvider:
"""Combine local model VAD with energy fallback for live microphone variance."""
def __init__(self, primary: Any, fallback: Any) -> None:
self.primary = primary
self.fallback = fallback
self.loaded = False
self._speech_ms = 0
self._silence_ms = 0
def load(self) -> None:
self.primary.load()
self.fallback.load()
self.loaded = True
def analyze(self, frame: AudioFrame) -> VadResult:
if not self.loaded:
raise ProviderError(
ErrorCode.VAD_MODEL_LOAD_FAILED,
"hybrid VAD provider is not loaded",
False,
"hybrid-vad",
"vad",
)
primary = self.primary.analyze(frame)
fallback = self.fallback.analyze(frame)
is_speech = primary.is_speech or fallback.is_speech
frame_ms = int(frame.metadata.get("duration_ms", 20))
if is_speech:
self._speech_ms += frame_ms
self._silence_ms = 0
else:
self._silence_ms += frame_ms
return VadResult(
is_speech=is_speech,
confidence=max(primary.confidence, fallback.confidence),
speech_ms=self._speech_ms,
silence_ms=self._silence_ms,
)
def reset(self) -> None:
self._speech_ms = 0
self._silence_ms = 0
self.primary.reset()
self.fallback.reset()
@dataclass(slots=True)
class VadRecorder:
provider: Any