[主说话人端点]:完成音色消失结束录音,包含临时音色画像、端点配置和回归测试

This commit is contained in:
mkbk
2026-06-17 21:33:33 +08:00
parent f9da304568
commit da91be6e5c
8 changed files with 372 additions and 12 deletions
+170
View File
@@ -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)