452 lines
16 KiB
Python
452 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult
|
|
from .speech_models import vad_model_path
|
|
|
|
|
|
class EnergyVadProvider:
|
|
def __init__(self, threshold: int = 500) -> 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
|
|
try:
|
|
import struct
|
|
|
|
sample_count = len(frame.pcm) // 2
|
|
if sample_count:
|
|
samples = struct.unpack("<" + "h" * sample_count, frame.pcm[: sample_count * 2])
|
|
return max(abs(sample) for sample in samples) > self.threshold
|
|
except Exception:
|
|
pass
|
|
return any(abs(byte - 128) > self.threshold for byte in frame.pcm)
|
|
|
|
|
|
class SherpaOnnxVadProvider:
|
|
def __init__(self, models_dir: str | Path, threshold: float = 0.5, sherpa_module: Any | None = None) -> None:
|
|
self.models_dir = Path(models_dir)
|
|
self.threshold = threshold
|
|
self.loaded = False
|
|
self._sherpa = sherpa_module
|
|
self._model: Any | None = None
|
|
self._speech_ms = 0
|
|
self._silence_ms = 0
|
|
|
|
def load(self) -> None:
|
|
model_path = vad_model_path(self.models_dir)
|
|
if not model_path.exists():
|
|
raise ProviderError(
|
|
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
|
f"sherpa-onnx VAD model path does not exist: {model_path}",
|
|
False,
|
|
"sherpa-onnx-vad",
|
|
"vad",
|
|
)
|
|
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.VAD_MODEL_LOAD_FAILED,
|
|
f"sherpa_onnx is not available: {exc}",
|
|
False,
|
|
"sherpa-onnx-vad",
|
|
"vad",
|
|
) from exc
|
|
try:
|
|
config = sherpa_onnx.VadModelConfig(
|
|
silero_vad=sherpa_onnx.SileroVadModelConfig(model=str(model_path), threshold=self.threshold),
|
|
sample_rate=16000,
|
|
)
|
|
self._model = sherpa_onnx.VadModel.create(config)
|
|
except Exception as exc:
|
|
raise ProviderError(
|
|
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
|
f"failed to load sherpa-onnx VAD model: {exc}",
|
|
False,
|
|
"sherpa-onnx-vad",
|
|
"vad",
|
|
) from exc
|
|
self.loaded = True
|
|
|
|
def analyze(self, frame: AudioFrame) -> VadResult:
|
|
if not self.loaded or self._model is None:
|
|
raise ProviderError(
|
|
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
|
"sherpa-onnx VAD provider is not loaded",
|
|
False,
|
|
"sherpa-onnx-vad",
|
|
"vad",
|
|
)
|
|
try:
|
|
import numpy as np
|
|
|
|
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
|
|
window_size = int(self._model.window_size())
|
|
if samples.size < window_size:
|
|
samples = np.pad(samples, (0, window_size - samples.size))
|
|
elif samples.size > window_size:
|
|
samples = samples[-window_size:]
|
|
is_speech = bool(self._model.is_speech(samples))
|
|
except Exception as exc:
|
|
raise ProviderError(
|
|
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
|
f"sherpa-onnx VAD analysis failed: {exc}",
|
|
True,
|
|
"sherpa-onnx-vad",
|
|
"vad",
|
|
) from exc
|
|
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
|
|
if self._model is not None:
|
|
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 SpeakerProfile:
|
|
vector: tuple[float, ...] | None = None
|
|
speaker_id: str | None = None
|
|
speech_ms: int = 0
|
|
|
|
|
|
@dataclass(slots=True)
|
|
class VadRecorder:
|
|
provider: Any
|
|
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
|
|
|
|
def finish(self, end_reason: str) -> AudioSegment:
|
|
return self._build_segment(end_reason)
|
|
|
|
|
|
@dataclass(slots=True)
|
|
class PrimarySpeakerVadRecorder(VadRecorder):
|
|
speaker_profile_ms: int = 600
|
|
speaker_profile_min_ms: int = 120
|
|
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:
|
|
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, self.speaker_profile_min_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)
|