[Wake/VAD/STT]:完成本地唤醒、人声端点检测与转写入口,包含小杰小杰唤醒、VAD 录音切分、Metadata STT 和 sherpa-onnx 错误边界
This commit is contained in:
@@ -18,11 +18,11 @@
|
|||||||
|
|
||||||
## 3. Wake/VAD/STT
|
## 3. Wake/VAD/STT
|
||||||
|
|
||||||
- [ ] 3.1 实现本地唤醒词 Provider;前置条件:Transport 测试路径可用;验收标准:支持“小杰小杰”关键词事件和置信度阈值;测试要点:命中、未命中、模型失败场景通过;优先级:P0;预计:60 分钟。
|
- [x] 3.1 实现本地唤醒词 Provider;前置条件:Transport 测试路径可用;验收标准:支持“小杰小杰”关键词事件和置信度阈值;测试要点:命中、未命中、模型失败场景通过;优先级:P0;预计:60 分钟。
|
||||||
- [ ] 3.2 实现 VAD Provider 和端点检测;前置条件:音频帧可回放;验收标准:检测说话开始、连续静音结束、无语音超时和最大录音保护;测试要点:四类 VAD 场景通过;优先级:P0;预计:60 分钟。
|
- [x] 3.2 实现 VAD Provider 和端点检测;前置条件:音频帧可回放;验收标准:检测说话开始、连续静音结束、无语音超时和最大录音保护;测试要点:四类 VAD 场景通过;优先级:P0;预计:60 分钟。
|
||||||
- [ ] 3.3 实现 STT Provider 协议、测试 Provider 和可选本地模型适配入口;前置条件:AudioSegment 可构造;验收标准:测试 Provider 能从 fixture metadata 转写,可选本地模型缺失时返回结构化错误;测试要点:成功、空文本、Provider 失败场景通过;优先级:P0;预计:60 分钟。
|
- [x] 3.3 实现 STT Provider 协议、测试 Provider 和可选本地模型适配入口;前置条件:AudioSegment 可构造;验收标准:测试 Provider 能从 fixture metadata 转写,可选本地模型缺失时返回结构化错误;测试要点:成功、空文本、Provider 失败场景通过;优先级:P0;预计:60 分钟。
|
||||||
- [ ] 3.4 实现 STT 文本验证规则;前置条件:STT Provider 已实现;验收标准:空文本、纯标点、过短音频不会进入 LLM;测试要点:文本验证和 pipeline 跳过 LLM 测试通过;优先级:P0;预计:45 分钟。
|
- [x] 3.4 实现 STT 文本验证规则;前置条件:STT Provider 已实现;验收标准:空文本、纯标点、过短音频不会进入 LLM;测试要点:文本验证和 pipeline 跳过 LLM 测试通过;优先级:P0;预计:45 分钟。
|
||||||
- [ ] 3.5 完成“Wake/VAD/STT”模块提交;前置条件:3.1 至 3.4 已完成;验收标准:先通过相关测试、compileall 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[Wake/VAD/STT]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。
|
- [x] 3.5 完成“Wake/VAD/STT”模块提交;前置条件:3.1 至 3.4 已完成;验收标准:先通过相关测试、compileall 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[Wake/VAD/STT]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。
|
||||||
|
|
||||||
## 4. LLM/TTS 与对话闭环
|
## 4. LLM/TTS 与对话闭环
|
||||||
|
|
||||||
|
|||||||
@@ -16,6 +16,9 @@ from .models import (
|
|||||||
WakeEvent,
|
WakeEvent,
|
||||||
)
|
)
|
||||||
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
|
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
|
||||||
|
from .wakeword import KeywordWakeWordProvider
|
||||||
|
from .vad import EnergyVadProvider, VadRecorder
|
||||||
|
from .stt import MetadataSttProvider, is_valid_transcript_text
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AppConfig",
|
"AppConfig",
|
||||||
@@ -25,6 +28,11 @@ __all__ = [
|
|||||||
"ErrorCode",
|
"ErrorCode",
|
||||||
"FileReplayTransport",
|
"FileReplayTransport",
|
||||||
"MemoryAudioTransport",
|
"MemoryAudioTransport",
|
||||||
|
"KeywordWakeWordProvider",
|
||||||
|
"EnergyVadProvider",
|
||||||
|
"VadRecorder",
|
||||||
|
"MetadataSttProvider",
|
||||||
|
"is_valid_transcript_text",
|
||||||
"Message",
|
"Message",
|
||||||
"PipelineState",
|
"PipelineState",
|
||||||
"PlaybackResult",
|
"PlaybackResult",
|
||||||
|
|||||||
@@ -0,0 +1,93 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from .models import AudioSegment, ErrorCode, ProviderError, Transcript
|
||||||
|
|
||||||
|
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
|
||||||
|
|
||||||
|
|
||||||
|
def is_valid_transcript_text(text: str) -> bool:
|
||||||
|
return bool(_MEANINGFUL_TEXT.search(text.strip()))
|
||||||
|
|
||||||
|
|
||||||
|
class MetadataSttProvider:
|
||||||
|
def __init__(self, language: str = "zh") -> None:
|
||||||
|
self.language = language
|
||||||
|
self.loaded = False
|
||||||
|
|
||||||
|
def load(self) -> None:
|
||||||
|
self.loaded = True
|
||||||
|
|
||||||
|
def transcribe(self, segment: AudioSegment) -> Transcript:
|
||||||
|
if not self.loaded:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||||
|
"STT provider is not loaded",
|
||||||
|
False,
|
||||||
|
"metadata-stt",
|
||||||
|
"stt",
|
||||||
|
)
|
||||||
|
text = str(segment.metadata.get("transcript", "")).strip()
|
||||||
|
if not is_valid_transcript_text(text):
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_EMPTY_TRANSCRIPT,
|
||||||
|
"STT produced no meaningful text",
|
||||||
|
True,
|
||||||
|
"metadata-stt",
|
||||||
|
"stt",
|
||||||
|
)
|
||||||
|
return Transcript(
|
||||||
|
text=text,
|
||||||
|
language=str(segment.metadata.get("language", self.language)),
|
||||||
|
confidence=float(segment.metadata.get("stt_confidence", 1.0)),
|
||||||
|
duration_ms=segment.duration_ms,
|
||||||
|
provider="metadata-stt",
|
||||||
|
raw_metadata=dict(segment.metadata),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SherpaOnnxSttProvider:
|
||||||
|
def __init__(self, model_path: str, language: str = "zh") -> None:
|
||||||
|
self.model_path = Path(model_path)
|
||||||
|
self.language = language
|
||||||
|
self.loaded = False
|
||||||
|
|
||||||
|
def load(self) -> None:
|
||||||
|
if not self.model_path.exists():
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_MODEL_MISSING,
|
||||||
|
f"sherpa-onnx STT model path does not exist: {self.model_path}",
|
||||||
|
False,
|
||||||
|
"sherpa-onnx-stt",
|
||||||
|
"stt",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
import sherpa_onnx # type: ignore[import-not-found] # noqa: F401
|
||||||
|
except Exception as exc:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||||
|
f"sherpa_onnx is not available: {exc}",
|
||||||
|
False,
|
||||||
|
"sherpa-onnx-stt",
|
||||||
|
"stt",
|
||||||
|
) from exc
|
||||||
|
self.loaded = True
|
||||||
|
|
||||||
|
def transcribe(self, segment: AudioSegment) -> Transcript:
|
||||||
|
if not self.loaded:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||||
|
"sherpa-onnx STT provider is not loaded",
|
||||||
|
False,
|
||||||
|
"sherpa-onnx-stt",
|
||||||
|
"stt",
|
||||||
|
)
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||||
|
"sherpa-onnx runtime transcription adapter requires a concrete model profile",
|
||||||
|
False,
|
||||||
|
"sherpa-onnx-stt",
|
||||||
|
"stt",
|
||||||
|
)
|
||||||
@@ -0,0 +1,123 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult
|
||||||
|
|
||||||
|
|
||||||
|
class EnergyVadProvider:
|
||||||
|
def __init__(self, threshold: int = 0) -> None:
|
||||||
|
self.threshold = threshold
|
||||||
|
self.loaded = False
|
||||||
|
self._speech_ms = 0
|
||||||
|
self._silence_ms = 0
|
||||||
|
|
||||||
|
def load(self) -> None:
|
||||||
|
self.loaded = True
|
||||||
|
|
||||||
|
def analyze(self, frame: AudioFrame) -> VadResult:
|
||||||
|
if not self.loaded:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
||||||
|
"VAD provider is not loaded",
|
||||||
|
False,
|
||||||
|
"energy-vad",
|
||||||
|
"vad",
|
||||||
|
)
|
||||||
|
is_speech = self._is_speech(frame)
|
||||||
|
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=0.9 if is_speech else 0.1,
|
||||||
|
speech_ms=self._speech_ms,
|
||||||
|
silence_ms=self._silence_ms,
|
||||||
|
)
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
self._speech_ms = 0
|
||||||
|
self._silence_ms = 0
|
||||||
|
|
||||||
|
def _is_speech(self, frame: AudioFrame) -> bool:
|
||||||
|
if "speech" in frame.metadata:
|
||||||
|
return bool(frame.metadata["speech"])
|
||||||
|
if not frame.pcm:
|
||||||
|
return False
|
||||||
|
return any(abs(byte - 128) > self.threshold for byte in frame.pcm)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class VadRecorder:
|
||||||
|
provider: EnergyVadProvider
|
||||||
|
min_duration_ms: int = 300
|
||||||
|
end_silence_ms: int = 200
|
||||||
|
no_speech_timeout_ms: int = 1000
|
||||||
|
max_recording_ms: int = 30000
|
||||||
|
started: bool = field(default=False, init=False)
|
||||||
|
frames: list[AudioFrame] = field(default_factory=list, init=False)
|
||||||
|
first_seen_ms: int | None = field(default=None, init=False)
|
||||||
|
start_time_ms: int | None = field(default=None, init=False)
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
self.started = False
|
||||||
|
self.frames: list[AudioFrame] = []
|
||||||
|
self.first_seen_ms: int | None = None
|
||||||
|
self.start_time_ms: int | None = None
|
||||||
|
|
||||||
|
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)
|
||||||
|
if result.is_speech:
|
||||||
|
if not self.started:
|
||||||
|
self.started = True
|
||||||
|
self.start_time_ms = frame.timestamp_ms
|
||||||
|
self.frames.append(frame)
|
||||||
|
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,
|
||||||
|
"energy-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 result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms:
|
||||||
|
return self._build_segment("silence")
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _build_segment(self, end_reason: str) -> AudioSegment:
|
||||||
|
if not self.frames:
|
||||||
|
raise ValueError("cannot build empty segment")
|
||||||
|
metadata: dict[str, object] = {"end_reason": end_reason}
|
||||||
|
for frame in self.frames:
|
||||||
|
metadata.update(dict(frame.metadata))
|
||||||
|
segment = AudioSegment(
|
||||||
|
pcm=b"".join(frame.pcm for frame in self.frames),
|
||||||
|
sample_rate=self.frames[0].sample_rate,
|
||||||
|
channels=self.frames[0].channels,
|
||||||
|
start_time_ms=self.frames[0].timestamp_ms,
|
||||||
|
end_time_ms=self.frames[-1].timestamp_ms
|
||||||
|
+ int(self.frames[-1].metadata.get("duration_ms", 20)),
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
self.reset()
|
||||||
|
self.provider.reset()
|
||||||
|
return segment
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent
|
||||||
|
|
||||||
|
|
||||||
|
class KeywordWakeWordProvider:
|
||||||
|
def __init__(self, keyword: str = "小杰小杰", threshold: float = 0.5) -> None:
|
||||||
|
self.keyword = keyword
|
||||||
|
self.threshold = threshold
|
||||||
|
self.loaded = False
|
||||||
|
|
||||||
|
def load(self) -> None:
|
||||||
|
self.loaded = True
|
||||||
|
|
||||||
|
def detect(self, frame: AudioFrame) -> WakeEvent | None:
|
||||||
|
if not self.loaded:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.WAKE_MODEL_LOAD_FAILED,
|
||||||
|
"wakeword provider is not loaded",
|
||||||
|
False,
|
||||||
|
"keyword-wakeword",
|
||||||
|
"wakeword",
|
||||||
|
)
|
||||||
|
metadata = frame.metadata
|
||||||
|
confidence = float(metadata.get("wake_confidence", 1.0 if metadata.get("wake") else 0.0))
|
||||||
|
phrase = str(metadata.get("wake_word", metadata.get("text", "")))
|
||||||
|
matched = bool(metadata.get("wake")) or phrase.strip() == self.keyword
|
||||||
|
if matched and confidence >= self.threshold:
|
||||||
|
return WakeEvent(self.keyword, confidence, frame.timestamp_ms)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class MissingWakeWordModelProvider:
|
||||||
|
def __init__(self, model_path: str) -> None:
|
||||||
|
self.model_path = model_path
|
||||||
|
|
||||||
|
def load(self) -> None:
|
||||||
|
raise ProviderError(
|
||||||
|
ErrorCode.WAKE_MODEL_MISSING,
|
||||||
|
f"wakeword model is missing: {self.model_path}",
|
||||||
|
False,
|
||||||
|
"wakeword-model",
|
||||||
|
"wakeword",
|
||||||
|
)
|
||||||
|
|
||||||
|
def detect(self, frame: AudioFrame) -> WakeEvent | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
return None
|
||||||
@@ -0,0 +1,107 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
|
||||||
|
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||||
|
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
|
||||||
|
from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider
|
||||||
|
|
||||||
|
|
||||||
|
def make_frame(
|
||||||
|
idx: int,
|
||||||
|
timestamp_ms: int,
|
||||||
|
*,
|
||||||
|
speech: bool = False,
|
||||||
|
metadata: dict[str, object] | None = None,
|
||||||
|
) -> AudioFrame:
|
||||||
|
data = b"\xff\xff" if speech else b"\x80\x80"
|
||||||
|
merged = {"speech": speech, "duration_ms": 20}
|
||||||
|
if metadata:
|
||||||
|
merged.update(metadata)
|
||||||
|
return AudioFrame(data, 16000, 1, timestamp_ms, idx, merged)
|
||||||
|
|
||||||
|
|
||||||
|
class WakeVadSttTests(unittest.TestCase):
|
||||||
|
def test_keyword_wakeword_detects_chinese_phrase(self) -> None:
|
||||||
|
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.7)
|
||||||
|
provider.load()
|
||||||
|
event = provider.detect(
|
||||||
|
make_frame(1, 100, metadata={"wake_word": "小杰小杰", "wake_confidence": 0.9})
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(event)
|
||||||
|
self.assertEqual(event.keyword, "小杰小杰")
|
||||||
|
|
||||||
|
def test_keyword_wakeword_ignores_low_confidence(self) -> None:
|
||||||
|
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.8)
|
||||||
|
provider.load()
|
||||||
|
self.assertIsNone(
|
||||||
|
provider.detect(make_frame(1, 100, metadata={"wake": True, "wake_confidence": 0.2}))
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_missing_wake_model_reports_structured_error(self) -> None:
|
||||||
|
with self.assertRaises(ProviderError) as raised:
|
||||||
|
MissingWakeWordModelProvider("/missing/model.onnx").load()
|
||||||
|
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
||||||
|
|
||||||
|
def test_vad_recorder_returns_segment_after_silence(self) -> None:
|
||||||
|
provider = EnergyVadProvider()
|
||||||
|
provider.load()
|
||||||
|
recorder = VadRecorder(provider, min_duration_ms=40, end_silence_ms=40)
|
||||||
|
frames = [
|
||||||
|
make_frame(1, 0, speech=True, metadata={"transcript": "你好"}),
|
||||||
|
make_frame(2, 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)
|
||||||
|
self.assertEqual(segment.metadata["end_reason"], "silence")
|
||||||
|
self.assertEqual(segment.metadata["transcript"], "你好")
|
||||||
|
|
||||||
|
def test_vad_recorder_returns_no_speech_timeout_error(self) -> None:
|
||||||
|
provider = EnergyVadProvider()
|
||||||
|
provider.load()
|
||||||
|
recorder = VadRecorder(provider, no_speech_timeout_ms=40)
|
||||||
|
result = None
|
||||||
|
for item in [make_frame(1, 0), make_frame(2, 20), make_frame(3, 40)]:
|
||||||
|
result = recorder.feed(item)
|
||||||
|
self.assertIsInstance(result, ProviderError)
|
||||||
|
self.assertEqual(result.code, ErrorCode.VAD_TIMEOUT_NO_SPEECH)
|
||||||
|
|
||||||
|
def test_metadata_stt_transcribes_fixture_text(self) -> None:
|
||||||
|
provider = MetadataSttProvider()
|
||||||
|
provider.load()
|
||||||
|
transcript = provider.transcribe(
|
||||||
|
AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "今天天气怎么样"})
|
||||||
|
)
|
||||||
|
self.assertEqual(transcript.normalized_text, "今天天气怎么样")
|
||||||
|
self.assertEqual(transcript.language, "zh")
|
||||||
|
|
||||||
|
def test_metadata_stt_rejects_empty_text(self) -> None:
|
||||||
|
provider = MetadataSttProvider()
|
||||||
|
provider.load()
|
||||||
|
with self.assertRaises(ProviderError) as raised:
|
||||||
|
provider.transcribe(AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "。!?"}))
|
||||||
|
self.assertEqual(raised.exception.code, ErrorCode.STT_EMPTY_TRANSCRIPT)
|
||||||
|
|
||||||
|
def test_transcript_validator_accepts_chinese_and_ascii(self) -> None:
|
||||||
|
self.assertTrue(is_valid_transcript_text("你好"))
|
||||||
|
self.assertTrue(is_valid_transcript_text("hello"))
|
||||||
|
self.assertFalse(is_valid_transcript_text("?! 。"))
|
||||||
|
|
||||||
|
def test_sherpa_stt_missing_model_is_structured(self) -> None:
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
provider = SherpaOnnxSttProvider(f"{tmp}/missing")
|
||||||
|
with self.assertRaises(ProviderError) as raised:
|
||||||
|
provider.load()
|
||||||
|
self.assertEqual(raised.exception.code, ErrorCode.STT_MODEL_MISSING)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user