[唤醒灵敏度与端点恢复]:完成唤醒提示顺序和Hybrid VAD优化,包含阈值默认值、缓冲时序和测试覆盖
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user