[主说话人端点]:完成音色消失结束录音,包含临时音色画像、端点配置和回归测试
This commit is contained in:
@@ -51,6 +51,6 @@
|
||||
- [x] 7.1 更新 OpenSpec 以描述 stage 化 pipeline、事件总线、TurnController 和主说话人端点;前置条件:公开参考已确认;验收标准:proposal/design/spec/tasks 覆盖新架构和任务;测试要点:OpenSpec strict;优先级:P0;预计:45 分钟。
|
||||
- [x] 7.2 实现 pipeline event bus 和终端事件映射;前置条件:7.1 完成;验收标准:所有 live 用户可见状态由事件产生;测试要点:事件顺序和终端文案测试;优先级:P0;预计:60 分钟。
|
||||
- [x] 7.3 实现 `TurnController` 和 `VoiceAssistantPipeline`;前置条件:7.2 完成;验收标准:`run-live` 使用统一 pipeline,成功/失败 turn 均恢复待机;测试要点:两轮 fake runtime、错误恢复、上下文回归;优先级:P0;预计:60 分钟。
|
||||
- [ ] 7.4 实现本轮主说话人端点;前置条件:7.3 完成;验收标准:主说话人音色消失约 300 ms 后结束采集;测试要点:一次提问后背景噪声不拖尾、短暂停顿不断句、画像不足回退;优先级:P0;预计:60 分钟。
|
||||
- [x] 7.4 实现本轮主说话人端点;前置条件:7.3 完成;验收标准:主说话人音色消失约 300 ms 后结束采集;测试要点:一次提问后背景噪声不拖尾、短暂停顿不断句、画像不足回退;优先级:P0;预计:60 分钟。
|
||||
- [ ] 7.5 更新 README、`.env.example`、本地 `.env` 非密钥配置;前置条件:7.2 至 7.4 完成;验收标准:运行说明匹配新 pipeline;测试要点:`--show-config` 不泄露 key;优先级:P0;预计:30 分钟。
|
||||
- [ ] 7.6 验证并提交“Pipeline 文档验收”模块;前置条件:7.1 至 7.5 完成;验收标准:compileall、unittest、security-check、model-check、device-check、OpenSpec strict 全通过;优先级:P0;预计:30 分钟。
|
||||
|
||||
@@ -19,7 +19,7 @@ from .models import (
|
||||
)
|
||||
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
|
||||
from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
|
||||
from .vad import EnergyVadProvider, HybridVadProvider, VadRecorder
|
||||
from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder
|
||||
from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from .conversation import ConversationContext
|
||||
from .llm import MockLlmProvider, OpenAICompatibleLlmProvider
|
||||
@@ -45,6 +45,7 @@ __all__ = [
|
||||
"SherpaOnnxKeywordWakeWordProvider",
|
||||
"EnergyVadProvider",
|
||||
"HybridVadProvider",
|
||||
"PrimarySpeakerVadRecorder",
|
||||
"VadRecorder",
|
||||
"CloudAsrSttProvider",
|
||||
"MetadataSttProvider",
|
||||
|
||||
@@ -26,6 +26,12 @@ class AppConfig:
|
||||
wake_kws_score: float = 1.0
|
||||
wake_ack_text: str = "我在"
|
||||
post_playback_drain_ms: int = 50
|
||||
pipeline_mode: str = "live_turn_based"
|
||||
endpoint_mode: str = "primary_speaker"
|
||||
speaker_profile_ms: int = 600
|
||||
speaker_absent_ms: int = 300
|
||||
speaker_similarity_threshold: float = 0.70
|
||||
speaker_min_rms: float = 0.012
|
||||
vad_provider: str = "hybrid"
|
||||
vad_threshold: float = 0.5
|
||||
vad_min_duration_ms: int = 250
|
||||
@@ -37,6 +43,7 @@ class AppConfig:
|
||||
tts_model: str = "mimo-v2.5-tts"
|
||||
tts_voice: str = "mimo_default"
|
||||
speech_models_dir: Path = Path("models")
|
||||
context_mode: str = "session_memory"
|
||||
context_max_messages: int = 12
|
||||
context_max_chars: int = 12000
|
||||
|
||||
@@ -67,6 +74,14 @@ class AppConfig:
|
||||
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", "50") or "50"),
|
||||
pipeline_mode=(get("PIPELINE_MODE", "live_turn_based") or "live_turn_based").lower(),
|
||||
endpoint_mode=(get("ENDPOINT_MODE", "primary_speaker") or "primary_speaker").lower(),
|
||||
speaker_profile_ms=int(get("SPEAKER_PROFILE_MS", "600") or "600"),
|
||||
speaker_absent_ms=int(get("SPEAKER_ABSENT_MS", "300") or "300"),
|
||||
speaker_similarity_threshold=float(
|
||||
get("SPEAKER_SIMILARITY_THRESHOLD", "0.70") or "0.70"
|
||||
),
|
||||
speaker_min_rms=float(get("SPEAKER_MIN_RMS", "0.012") or "0.012"),
|
||||
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"),
|
||||
@@ -78,6 +93,7 @@ class AppConfig:
|
||||
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
|
||||
tts_voice=get("TTS_VOICE", "mimo_default") or "mimo_default",
|
||||
speech_models_dir=Path(get("SPEECH_MODELS_DIR", "models") or "models"),
|
||||
context_mode=(get("CONTEXT_MODE", "session_memory") or "session_memory").lower(),
|
||||
context_max_messages=int(get("CONTEXT_MAX_MESSAGES", "12") or "12"),
|
||||
context_max_chars=int(get("CONTEXT_MAX_CHARS", "12000") or "12000"),
|
||||
)
|
||||
@@ -178,6 +194,60 @@ class AppConfig:
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.pipeline_mode not in {"live_turn_based"}:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_PIPELINE_MODE must be live_turn_based",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.endpoint_mode not in {"primary_speaker", "vad"}:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_ENDPOINT_MODE must be primary_speaker or vad",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
for name, value in {
|
||||
"OWNER_SPEAKER_PROFILE_MS": self.speaker_profile_ms,
|
||||
"OWNER_SPEAKER_ABSENT_MS": self.speaker_absent_ms,
|
||||
}.items():
|
||||
if value <= 0:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
f"{name} must be positive",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if not 0 < self.speaker_similarity_threshold <= 1:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_SPEAKER_SIMILARITY_THRESHOLD must be in (0, 1]",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.speaker_min_rms <= 0:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_SPEAKER_MIN_RMS must be positive",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.vad_provider not in {"hybrid", "local", "energy"}:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
@@ -224,6 +294,16 @@ class AppConfig:
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
if self.context_mode not in {"session_memory"}:
|
||||
errors.append(
|
||||
ProviderError(
|
||||
ErrorCode.CONFIG_MISSING_VALUE,
|
||||
"OWNER_CONTEXT_MODE must be session_memory",
|
||||
False,
|
||||
"config",
|
||||
"startup",
|
||||
)
|
||||
)
|
||||
return errors
|
||||
|
||||
def api_url(self, path: str) -> str:
|
||||
|
||||
@@ -33,7 +33,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, HybridVadProvider, SherpaOnnxVadProvider, VadRecorder
|
||||
from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, SherpaOnnxVadProvider, VadRecorder
|
||||
from .wakeword import SherpaOnnxKeywordWakeWordProvider
|
||||
|
||||
|
||||
@@ -345,6 +345,23 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
|
||||
vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold)
|
||||
else:
|
||||
vad_provider = EnergyVadProvider()
|
||||
recorder_cls = PrimarySpeakerVadRecorder if config.endpoint_mode == "primary_speaker" else VadRecorder
|
||||
recorder_kwargs = {
|
||||
"provider": vad_provider,
|
||||
"min_duration_ms": config.vad_min_duration_ms,
|
||||
"end_silence_ms": config.vad_end_silence_ms,
|
||||
"no_speech_timeout_ms": config.vad_no_speech_timeout_ms,
|
||||
"max_recording_ms": config.vad_max_recording_ms,
|
||||
}
|
||||
if recorder_cls is PrimarySpeakerVadRecorder:
|
||||
recorder_kwargs.update(
|
||||
{
|
||||
"speaker_profile_ms": config.speaker_profile_ms,
|
||||
"speaker_absent_ms": config.speaker_absent_ms,
|
||||
"similarity_threshold": config.speaker_similarity_threshold,
|
||||
"min_rms": config.speaker_min_rms,
|
||||
}
|
||||
)
|
||||
return VoiceAssistantPipeline(
|
||||
config=config,
|
||||
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
|
||||
@@ -355,13 +372,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
|
||||
threshold=config.wake_kws_threshold,
|
||||
score=config.wake_kws_score,
|
||||
),
|
||||
vad_recorder=VadRecorder(
|
||||
vad_provider,
|
||||
min_duration_ms=config.vad_min_duration_ms,
|
||||
end_silence_ms=config.vad_end_silence_ms,
|
||||
no_speech_timeout_ms=config.vad_no_speech_timeout_ms,
|
||||
max_recording_ms=config.vad_max_recording_ms,
|
||||
),
|
||||
vad_recorder=recorder_cls(**recorder_kwargs),
|
||||
stt=stt,
|
||||
llm=OpenAICompatibleLlmProvider(config),
|
||||
tts=tts,
|
||||
|
||||
@@ -204,6 +204,13 @@ class HybridVadProvider:
|
||||
self.fallback.reset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SpeakerProfile:
|
||||
vector: tuple[float, ...] | None = None
|
||||
speaker_id: str | None = None
|
||||
speech_ms: int = 0
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VadRecorder:
|
||||
provider: Any
|
||||
@@ -275,3 +282,166 @@ class VadRecorder:
|
||||
self.reset()
|
||||
self.provider.reset()
|
||||
return segment
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PrimarySpeakerVadRecorder(VadRecorder):
|
||||
speaker_profile_ms: int = 600
|
||||
speaker_absent_ms: int = 300
|
||||
similarity_threshold: float = 0.70
|
||||
min_rms: float = 0.012
|
||||
profile: SpeakerProfile = field(default_factory=SpeakerProfile, init=False)
|
||||
profile_vectors: list[tuple[float, ...]] = field(default_factory=list, init=False)
|
||||
primary_absent_ms: int = field(default=0, init=False)
|
||||
|
||||
def reset(self) -> None:
|
||||
VadRecorder.reset(self)
|
||||
self.profile = SpeakerProfile()
|
||||
self.profile_vectors = []
|
||||
self.primary_absent_ms = 0
|
||||
|
||||
def feed(self, frame: AudioFrame) -> AudioSegment | ProviderError | None:
|
||||
if self.first_seen_ms is None:
|
||||
self.first_seen_ms = frame.timestamp_ms
|
||||
result = self.provider.analyze(frame)
|
||||
frame_ms = int(frame.metadata.get("duration_ms", 20))
|
||||
|
||||
if result.is_speech:
|
||||
if not self.started:
|
||||
self.started = True
|
||||
self.start_time_ms = frame.timestamp_ms
|
||||
self.frames.append(frame)
|
||||
self._update_profile(frame, frame_ms)
|
||||
elif self.started:
|
||||
self.frames.append(frame)
|
||||
|
||||
if not self.started:
|
||||
elapsed = frame.timestamp_ms - self.first_seen_ms
|
||||
if elapsed >= self.no_speech_timeout_ms:
|
||||
return ProviderError(
|
||||
ErrorCode.VAD_TIMEOUT_NO_SPEECH,
|
||||
"no speech detected after wakeword",
|
||||
True,
|
||||
"primary-speaker-vad",
|
||||
"vad",
|
||||
)
|
||||
return None
|
||||
|
||||
start_time = self.start_time_ms if self.start_time_ms is not None else frame.timestamp_ms
|
||||
duration = frame.timestamp_ms - start_time
|
||||
if duration >= self.max_recording_ms:
|
||||
return self._build_segment("max_recording")
|
||||
if self._profile_ready():
|
||||
if self._matches_primary(frame):
|
||||
self.primary_absent_ms = 0
|
||||
else:
|
||||
self.primary_absent_ms += frame_ms
|
||||
if self.primary_absent_ms >= self.speaker_absent_ms and duration >= self.min_duration_ms:
|
||||
return self._build_segment("primary_speaker_absent")
|
||||
if result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms:
|
||||
return self._build_segment("silence")
|
||||
return None
|
||||
|
||||
def _update_profile(self, frame: AudioFrame, frame_ms: int) -> None:
|
||||
if self.profile.speech_ms >= self.speaker_profile_ms:
|
||||
return
|
||||
speaker_id = frame.metadata.get("speaker_id")
|
||||
if speaker_id is not None:
|
||||
if self.profile.speaker_id is None:
|
||||
self.profile.speaker_id = str(speaker_id)
|
||||
if str(speaker_id) == self.profile.speaker_id:
|
||||
self.profile.speech_ms += frame_ms
|
||||
return
|
||||
vector = extract_timbre_vector(frame, min_rms=self.min_rms)
|
||||
if vector is None:
|
||||
return
|
||||
self.profile_vectors.append(vector)
|
||||
self.profile.speech_ms += frame_ms
|
||||
if self.profile_vectors:
|
||||
width = len(self.profile_vectors[0])
|
||||
averaged = []
|
||||
for index in range(width):
|
||||
averaged.append(sum(item[index] for item in self.profile_vectors) / len(self.profile_vectors))
|
||||
self.profile.vector = tuple(averaged)
|
||||
|
||||
def _profile_ready(self) -> bool:
|
||||
minimum_ms = min(self.speaker_profile_ms, max(120, self.min_duration_ms))
|
||||
return self.profile.speech_ms >= minimum_ms and (
|
||||
self.profile.speaker_id is not None or self.profile.vector is not None
|
||||
)
|
||||
|
||||
def _matches_primary(self, frame: AudioFrame) -> bool:
|
||||
speaker_id = frame.metadata.get("speaker_id")
|
||||
if self.profile.speaker_id is not None:
|
||||
return speaker_id is not None and str(speaker_id) == self.profile.speaker_id
|
||||
if self.profile.vector is None:
|
||||
return True
|
||||
vector = extract_timbre_vector(frame, min_rms=self.min_rms)
|
||||
if vector is None:
|
||||
return False
|
||||
return cosine_similarity(self.profile.vector, vector) >= self.similarity_threshold
|
||||
|
||||
|
||||
def extract_timbre_vector(frame: AudioFrame, *, min_rms: float) -> tuple[float, ...] | None:
|
||||
if "timbre_vector" in frame.metadata:
|
||||
raw = frame.metadata["timbre_vector"]
|
||||
if isinstance(raw, (list, tuple)) and raw:
|
||||
return tuple(float(item) for item in raw)
|
||||
if not frame.pcm:
|
||||
return None
|
||||
try:
|
||||
import numpy as np
|
||||
|
||||
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32)
|
||||
if samples.size == 0:
|
||||
return None
|
||||
if frame.channels > 1:
|
||||
samples = samples.reshape(-1, frame.channels).mean(axis=1)
|
||||
normalized = samples / 32768.0
|
||||
rms = float(np.sqrt(np.mean(normalized * normalized)))
|
||||
if rms < min_rms:
|
||||
return None
|
||||
signs = np.signbit(normalized)
|
||||
zcr = float(np.mean(signs[1:] != signs[:-1])) if normalized.size > 1 else 0.0
|
||||
windowed = normalized * np.hanning(normalized.size)
|
||||
spectrum = np.abs(np.fft.rfft(windowed))
|
||||
total = float(np.sum(spectrum))
|
||||
if total <= 1e-9:
|
||||
return None
|
||||
freqs = np.fft.rfftfreq(normalized.size, 1.0 / frame.sample_rate)
|
||||
nyquist = max(frame.sample_rate / 2.0, 1.0)
|
||||
centroid = float(np.sum(freqs * spectrum) / total) / nyquist
|
||||
bandwidth = float(np.sqrt(np.sum(((freqs / nyquist - centroid) ** 2) * spectrum) / total))
|
||||
cumulative = np.cumsum(spectrum)
|
||||
rolloff_index = int(np.searchsorted(cumulative, 0.85 * cumulative[-1]))
|
||||
rolloff = float(freqs[min(rolloff_index, freqs.size - 1)] / nyquist)
|
||||
flatness = float(np.exp(np.mean(np.log(spectrum + 1e-9))) / (np.mean(spectrum) + 1e-9))
|
||||
|
||||
def band_ratio(low: float, high: float) -> float:
|
||||
mask = (freqs >= low) & (freqs < high)
|
||||
return float(np.sum(spectrum[mask]) / total)
|
||||
|
||||
return (
|
||||
rms,
|
||||
zcr,
|
||||
centroid,
|
||||
bandwidth,
|
||||
rolloff,
|
||||
flatness,
|
||||
band_ratio(80, 500),
|
||||
band_ratio(500, 2000),
|
||||
band_ratio(2000, nyquist),
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def cosine_similarity(left: tuple[float, ...], right: tuple[float, ...]) -> float:
|
||||
if len(left) != len(right):
|
||||
return 0.0
|
||||
numerator = sum(a * b for a, b in zip(left, right))
|
||||
left_norm = sum(a * a for a in left) ** 0.5
|
||||
right_norm = sum(b * b for b in right) ** 0.5
|
||||
if left_norm <= 1e-9 or right_norm <= 1e-9:
|
||||
return 0.0
|
||||
return numerator / (left_norm * right_norm)
|
||||
|
||||
@@ -22,9 +22,10 @@ from owner_voice_pet.events import (
|
||||
from owner_voice_pet.llm import MockLlmProvider
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
|
||||
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
|
||||
from owner_voice_pet.runtime import build_live_runtime
|
||||
from owner_voice_pet.transport import MemoryAudioTransport
|
||||
from owner_voice_pet.tts import SineTtsProvider
|
||||
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
|
||||
from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder
|
||||
from owner_voice_pet.wakeword import KeywordWakeWordProvider
|
||||
|
||||
|
||||
@@ -181,6 +182,10 @@ class LiveRuntimeTests(unittest.TestCase):
|
||||
self.assertEqual(llm.calls[0][-1].content, "第一问")
|
||||
self.assertNotIn("小杰小杰", llm.calls[0][-1].content)
|
||||
|
||||
def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None:
|
||||
runtime = build_live_runtime(AppConfig(llm_api_key="secret"))
|
||||
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -74,6 +74,12 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
self.assertEqual(config.wake_kws_score, 1.0)
|
||||
self.assertEqual(config.wake_ack_text, "我在")
|
||||
self.assertEqual(config.post_playback_drain_ms, 50)
|
||||
self.assertEqual(config.pipeline_mode, "live_turn_based")
|
||||
self.assertEqual(config.endpoint_mode, "primary_speaker")
|
||||
self.assertEqual(config.speaker_profile_ms, 600)
|
||||
self.assertEqual(config.speaker_absent_ms, 300)
|
||||
self.assertEqual(config.speaker_similarity_threshold, 0.70)
|
||||
self.assertEqual(config.speaker_min_rms, 0.012)
|
||||
self.assertEqual(config.vad_provider, "hybrid")
|
||||
self.assertEqual(config.vad_threshold, 0.5)
|
||||
self.assertEqual(config.vad_min_duration_ms, 250)
|
||||
@@ -85,6 +91,7 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
self.assertEqual(config.tts_model, "mimo-v2.5-tts")
|
||||
self.assertEqual(config.tts_voice, "mimo_default")
|
||||
self.assertEqual(str(config.speech_models_dir), "models")
|
||||
self.assertEqual(config.context_mode, "session_memory")
|
||||
self.assertTrue(config.llm_stream)
|
||||
self.assertEqual(config.validate_basic(), [])
|
||||
|
||||
@@ -103,6 +110,16 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_VAD_PROVIDER" in error.message for error in errors))
|
||||
|
||||
def test_endpoint_provider_must_be_primary_speaker_or_vad(self) -> None:
|
||||
config = AppConfig(endpoint_mode="invalid")
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_ENDPOINT_MODE" in error.message for error in errors))
|
||||
|
||||
def test_speaker_similarity_threshold_range_is_validated(self) -> None:
|
||||
config = AppConfig(speaker_similarity_threshold=1.5)
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_SPEAKER_SIMILARITY_THRESHOLD" in error.message for error in errors))
|
||||
|
||||
def test_missing_dotenv_uses_non_secret_defaults(self) -> None:
|
||||
config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env")
|
||||
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
|
||||
|
||||
@@ -7,7 +7,7 @@ import unittest
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
|
||||
from owner_voice_pet.config import AppConfig
|
||||
from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, VadRecorder
|
||||
from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder
|
||||
from owner_voice_pet.wakeword import (
|
||||
KeywordWakeWordProvider,
|
||||
MissingWakeWordModelProvider,
|
||||
@@ -111,6 +111,82 @@ class WakeVadSttTests(unittest.TestCase):
|
||||
self.assertTrue(result.is_speech)
|
||||
self.assertEqual(result.speech_ms, 20)
|
||||
|
||||
def test_primary_speaker_endpoint_stops_on_background_noise(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
recorder = PrimarySpeakerVadRecorder(
|
||||
provider,
|
||||
min_duration_ms=40,
|
||||
end_silence_ms=1000,
|
||||
speaker_profile_ms=40,
|
||||
speaker_absent_ms=40,
|
||||
)
|
||||
frames = [
|
||||
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁"}),
|
||||
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
|
||||
make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}),
|
||||
make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}),
|
||||
make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁第二次"}),
|
||||
]
|
||||
segment = None
|
||||
consumed = 0
|
||||
for item in frames:
|
||||
consumed += 1
|
||||
result = recorder.feed(item)
|
||||
if isinstance(result, AudioSegment):
|
||||
segment = result
|
||||
break
|
||||
self.assertIsNotNone(segment)
|
||||
assert segment is not None
|
||||
self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent")
|
||||
self.assertEqual(consumed, 4)
|
||||
self.assertEqual(segment.metadata["transcript"], "你是谁")
|
||||
|
||||
def test_primary_speaker_endpoint_allows_short_pause(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
recorder = PrimarySpeakerVadRecorder(
|
||||
provider,
|
||||
min_duration_ms=40,
|
||||
end_silence_ms=1000,
|
||||
speaker_profile_ms=40,
|
||||
speaker_absent_ms=60,
|
||||
)
|
||||
frames = [
|
||||
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner"}),
|
||||
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
|
||||
make_frame(3, 40, speech=False),
|
||||
make_frame(4, 60, speech=True, metadata={"speaker_id": "owner"}),
|
||||
make_frame(5, 80, speech=False),
|
||||
make_frame(6, 100, speech=False),
|
||||
make_frame(7, 120, speech=False),
|
||||
]
|
||||
results = [recorder.feed(item) for item in frames]
|
||||
self.assertIsNone(results[2])
|
||||
self.assertIsNone(results[3])
|
||||
self.assertIsInstance(results[-1], AudioSegment)
|
||||
assert isinstance(results[-1], AudioSegment)
|
||||
self.assertEqual(results[-1].metadata["end_reason"], "primary_speaker_absent")
|
||||
|
||||
def test_primary_speaker_endpoint_falls_back_to_vad_when_profile_missing(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
recorder = PrimarySpeakerVadRecorder(provider, min_duration_ms=40, end_silence_ms=40)
|
||||
frames = [
|
||||
AudioFrame(b"", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "transcript": "你好"}),
|
||||
AudioFrame(b"", 16000, 1, 20, 2, {"duration_ms": 20, "speech": True}),
|
||||
make_frame(3, 40, speech=False),
|
||||
make_frame(4, 60, speech=False),
|
||||
]
|
||||
segment = None
|
||||
for item in frames:
|
||||
result = recorder.feed(item)
|
||||
if isinstance(result, AudioSegment):
|
||||
segment = result
|
||||
self.assertIsNotNone(segment)
|
||||
assert segment is not None
|
||||
self.assertEqual(segment.metadata["end_reason"], "silence")
|
||||
|
||||
def test_metadata_stt_transcribes_fixture_text(self) -> None:
|
||||
provider = MetadataSttProvider()
|
||||
provider.load()
|
||||
|
||||
Reference in New Issue
Block a user