From 1515f2e167038249162286d299a18f2c02f5e813 Mon Sep 17 00:00:00 2001 From: Yoann Date: Fri, 31 Jul 2026 16:02:48 +0200 Subject: [PATCH 1/5] feat: speaker diarization with voice-print identification Add optional speaker diarization via sherpa-onnx (pyannote segmentation + wespeaker embeddings, ONNX on CPU, fully local) behind a Diarizer port, plus voice-print enrolment so SPEAKER_00 labels become real names. - domain stays pure: SpeakerTurn, VoicePrint, assign_speakers, rename_speakers, cosine_similarity/match_speakers, longest_turn_per_speaker are plain Python with no numpy and no sherpa - sherpa_onnx is imported in adapters/ only, so the engine stays swappable through the port - --diarize turns off silence removal and denoising: silenceremove shifts the timeline away from the transcript, and dynaudnorm/afftdn degrade the speaker embeddings - diarization always runs on the exact file that was transcribed - vox speakers add/list stores voice prints in ~/.vox/voiceprints.json, --identify then renames labels on every later video - outputs carry the speaker: [NAME] in SRT, per-speaker blocks in TXT, "speaker" field in JSON - worker threads default to cpu_count - 2, floored at 2 Verified end to end on a two-voice dialogue: turns alternate correctly, enrolment then identification maps both speakers to their real names. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Uz3rnGtYHs6YcBNr6MkY1z --- pyproject.toml | 2 + src/vox/adapters/audio_decoding.py | 41 ++++++ src/vox/adapters/cli/app.py | 2 + src/vox/adapters/cli/speakers_cmd.py | 65 +++++++++ src/vox/adapters/cli/transcribe_cmd.py | 27 ++++ src/vox/adapters/cpu_threads.py | 14 ++ src/vox/adapters/diarization_models.py | 25 ++++ src/vox/adapters/disk_file_writer.py | 30 ++++- src/vox/adapters/json_voice_print_store.py | 24 ++++ src/vox/adapters/sherpa_diarizer.py | 75 +++++++++++ .../adapters/sherpa_voice_print_extractor.py | 43 ++++++ src/vox/adapters/system_dep_checker.py | 1 + src/vox/models/exceptions.py | 4 + src/vox/models/segment.py | 1 + src/vox/models/speaker_assignment.py | 31 +++++ src/vox/models/speaker_renaming.py | 18 +++ src/vox/models/speaker_turn.py | 18 +++ src/vox/models/turn_selection.py | 16 +++ src/vox/models/voice_matching.py | 43 ++++++ src/vox/models/voice_print.py | 15 +++ src/vox/ports/diarizer.py | 12 ++ src/vox/ports/voice_print_extractor.py | 11 ++ src/vox/ports/voice_print_store.py | 9 ++ src/vox/schemas/transcribe.json | 14 ++ src/vox/use_cases/enroll_speaker.py | 53 ++++++++ src/vox/use_cases/identify_speakers.py | 42 ++++++ src/vox/use_cases/transcribe.py | 56 ++++++-- tests/fakes/fake_diarizer.py | 21 +++ tests/fakes/fake_speaker_identifier.py | 17 +++ tests/fakes/fake_voice_print_extractor.py | 18 +++ tests/fakes/fake_voice_print_store.py | 15 +++ tests/unit/adapters/test_cpu_threads.py | 24 ++++ tests/unit/adapters/test_disk_file_writer.py | 78 +++++++++++ .../adapters/test_json_voice_print_store.py | 37 +++++ tests/unit/adapters/test_sherpa_diarizer.py | 38 ++++++ tests/unit/models/test_segment.py | 10 ++ tests/unit/models/test_speaker_assignment.py | 72 ++++++++++ tests/unit/models/test_speaker_renaming.py | 39 ++++++ tests/unit/models/test_speaker_turn.py | 29 ++++ tests/unit/models/test_turn_selection.py | 29 ++++ tests/unit/models/test_voice_matching.py | 54 ++++++++ tests/unit/models/test_voice_print.py | 20 +++ tests/unit/use_cases/test_enroll_speaker.py | 62 +++++++++ .../unit/use_cases/test_identify_speakers.py | 60 +++++++++ tests/unit/use_cases/test_transcribe.py | 126 ++++++++++++++++++ uv.lock | 43 ++++++ 46 files changed, 1472 insertions(+), 12 deletions(-) create mode 100644 src/vox/adapters/audio_decoding.py create mode 100644 src/vox/adapters/cli/speakers_cmd.py create mode 100644 src/vox/adapters/cpu_threads.py create mode 100644 src/vox/adapters/diarization_models.py create mode 100644 src/vox/adapters/json_voice_print_store.py create mode 100644 src/vox/adapters/sherpa_diarizer.py create mode 100644 src/vox/adapters/sherpa_voice_print_extractor.py create mode 100644 src/vox/models/speaker_assignment.py create mode 100644 src/vox/models/speaker_renaming.py create mode 100644 src/vox/models/speaker_turn.py create mode 100644 src/vox/models/turn_selection.py create mode 100644 src/vox/models/voice_matching.py create mode 100644 src/vox/models/voice_print.py create mode 100644 src/vox/ports/diarizer.py create mode 100644 src/vox/ports/voice_print_extractor.py create mode 100644 src/vox/ports/voice_print_store.py create mode 100644 src/vox/use_cases/enroll_speaker.py create mode 100644 src/vox/use_cases/identify_speakers.py create mode 100644 tests/fakes/fake_diarizer.py create mode 100644 tests/fakes/fake_speaker_identifier.py create mode 100644 tests/fakes/fake_voice_print_extractor.py create mode 100644 tests/fakes/fake_voice_print_store.py create mode 100644 tests/unit/adapters/test_cpu_threads.py create mode 100644 tests/unit/adapters/test_json_voice_print_store.py create mode 100644 tests/unit/adapters/test_sherpa_diarizer.py create mode 100644 tests/unit/models/test_speaker_assignment.py create mode 100644 tests/unit/models/test_speaker_renaming.py create mode 100644 tests/unit/models/test_speaker_turn.py create mode 100644 tests/unit/models/test_turn_selection.py create mode 100644 tests/unit/models/test_voice_matching.py create mode 100644 tests/unit/models/test_voice_print.py create mode 100644 tests/unit/use_cases/test_enroll_speaker.py create mode 100644 tests/unit/use_cases/test_identify_speakers.py diff --git a/pyproject.toml b/pyproject.toml index 79af143..2548456 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,6 +28,8 @@ classifiers = [ dependencies = [ "click>=8.1", "mlx-whisper>=0.4", + "sherpa-onnx>=1.13.4", + "sherpa-onnx-core>=1.13.4", "yt-dlp>=2024.0.0", ] diff --git a/src/vox/adapters/audio_decoding.py b/src/vox/adapters/audio_decoding.py new file mode 100644 index 0000000..7e69fbd --- /dev/null +++ b/src/vox/adapters/audio_decoding.py @@ -0,0 +1,41 @@ +import subprocess +from pathlib import Path + +from vox.models.exceptions import DiarizationError + +SAMPLE_RATE = 16000 + + +def decode_to_mono_16k(audio_path: Path): + import numpy as np + + result = subprocess.run( + _ffmpeg_decode_command(audio_path), + capture_output=True, + check=False, + ) + if result.returncode != 0: + stderr = result.stderr.decode(errors="replace") + raise DiarizationError(f"ffmpeg failed to decode audio: {stderr}") + return np.frombuffer(result.stdout, dtype=np.float32) + + +def slice_samples(samples, start: float, end: float): + return samples[int(start * SAMPLE_RATE) : int(end * SAMPLE_RATE)] + + +def _ffmpeg_decode_command(audio_path: Path) -> list[str]: + return [ + "ffmpeg", + "-v", + "quiet", + "-i", + str(audio_path), + "-ar", + str(SAMPLE_RATE), + "-ac", + "1", + "-f", + "f32le", + "-", + ] diff --git a/src/vox/adapters/cli/app.py b/src/vox/adapters/cli/app.py index 2cb2d57..9cfd8e2 100644 --- a/src/vox/adapters/cli/app.py +++ b/src/vox/adapters/cli/app.py @@ -7,6 +7,7 @@ from vox.adapters.cli.init_cmd import init from vox.adapters.cli.models_cmd import models from vox.adapters.cli.schema_cmd import schema +from vox.adapters.cli.speakers_cmd import speakers from vox.adapters.cli.transcribe_cmd import transcribe @@ -29,6 +30,7 @@ def main() -> None: main.add_command(doctor) main.add_command(schema) main.add_command(models) +main.add_command(speakers) def _is_agent_mode() -> bool: diff --git a/src/vox/adapters/cli/speakers_cmd.py b/src/vox/adapters/cli/speakers_cmd.py new file mode 100644 index 0000000..73ec687 --- /dev/null +++ b/src/vox/adapters/cli/speakers_cmd.py @@ -0,0 +1,65 @@ +import json as json_mod +import sys + +import click + +from vox.adapters.json_voice_print_store import JsonVoicePrintStore +from vox.adapters.sherpa_voice_print_extractor import SherpaVoicePrintExtractor +from vox.models.exceptions import VoxError +from vox.use_cases.enroll_speaker import EnrollSpeakerRequest, EnrollSpeakerUseCase + + +@click.group() +def speakers(): + """Manage known voices used by --identify.""" + + +@speakers.command() +@click.argument("name") +@click.option("--from", "source", required=True, help="Audio/video file to sample") +@click.option( + "--at", + required=True, + help="Time span of clean speech, in seconds: START-END (e.g. 120-140)", +) +def add(name, source, at): + """Register NAME's voice from a span of audio.""" + try: + start, end = _parse_span(at) + response = EnrollSpeakerUseCase( + extractor=SherpaVoicePrintExtractor(), + store=JsonVoicePrintStore(), + ).execute( + EnrollSpeakerRequest( + name=name, + audio_path=source, + start=start, + end=end, + ) + ) + except VoxError as e: + click.echo(f"Error: {e}", err=True) + sys.exit(1) + click.echo( + json_mod.dumps( + {"name": response.name, "seconds_used": response.seconds_used}, + indent=2, + ) + ) + + +@speakers.command(name="list") +def list_speakers(): + """List registered voices.""" + prints = JsonVoicePrintStore().load_all() + click.echo(json_mod.dumps({"speakers": [p.name for p in prints]}, indent=2)) + + +def _parse_span(raw: str) -> tuple[float, float]: + parts = raw.split("-") + if len(parts) != 2: + raise VoxError(f"invalid span '{raw}', expected START-END (e.g. 120-140)") + try: + return float(parts[0]), float(parts[1]) + except ValueError: + raise VoxError(f"invalid span '{raw}', expected numbers in seconds") from None diff --git a/src/vox/adapters/cli/transcribe_cmd.py b/src/vox/adapters/cli/transcribe_cmd.py index c568fb3..32da7db 100644 --- a/src/vox/adapters/cli/transcribe_cmd.py +++ b/src/vox/adapters/cli/transcribe_cmd.py @@ -8,14 +8,18 @@ from vox.adapters.click_progress import ClickProgressReporter from vox.adapters.disk_file_writer import DiskFileWriter from vox.adapters.ffmpeg_audio_cleaner import FfmpegAudioCleaner +from vox.adapters.json_voice_print_store import JsonVoicePrintStore from vox.adapters.mlx_transcriber import MlxTranscriber from vox.adapters.openai_transcriber import OpenAITranscriber +from vox.adapters.sherpa_diarizer import SherpaDiarizer +from vox.adapters.sherpa_voice_print_extractor import SherpaVoicePrintExtractor from vox.adapters.ytdlp_downloader import YtdlpDownloader from vox.models.exceptions import VoxError from vox.models.openai_model import OpenAIModel from vox.models.transcription_backend import TranscriptionBackend from vox.models.whisper_model import WhisperModel from vox.ports.transcriber import Transcriber +from vox.use_cases.identify_speakers import IdentifySpeakersUseCase from vox.use_cases.transcribe import TranscribeRequest, TranscribeUseCase @@ -44,6 +48,18 @@ default="local", help="local (MLX, default) | openai (cloud API)", ) +@click.option("--diarize", is_flag=True, help="Identify speakers (who said what)") +@click.option( + "--identify", + is_flag=True, + help="Replace SPEAKER_xx by known names (see: vox speakers add)", +) +@click.option( + "--speakers", + type=int, + default=None, + help="Known speaker count (auto-detected if omitted)", +) def transcribe( source, language, @@ -59,6 +75,9 @@ def transcribe( no_cookies, browser, backend, + diarize, + identify, + speakers, ): source, language, model = _apply_json_overrides( json_payload, source, language, model @@ -80,6 +99,9 @@ def transcribe( no_clean=no_clean, no_download=no_download, dry_run=dry_run, + diarize=diarize, + identify=identify, + num_speakers=speakers, ) try: response = use_case.execute(request) @@ -121,6 +143,11 @@ def _build_use_case( transcriber=_build_transcriber(backend), file_writer=DiskFileWriter(), progress=ClickProgressReporter(), + diarizer=SherpaDiarizer(), + speaker_identifier=IdentifySpeakersUseCase( + extractor=SherpaVoicePrintExtractor(), + store=JsonVoicePrintStore(), + ), ) diff --git a/src/vox/adapters/cpu_threads.py b/src/vox/adapters/cpu_threads.py new file mode 100644 index 0000000..f38f2ef --- /dev/null +++ b/src/vox/adapters/cpu_threads.py @@ -0,0 +1,14 @@ +import os + +_SPARE_CORES = 2 +_MIN_THREADS = 2 + + +def default_num_threads() -> int: + return choose_num_threads(os.cpu_count()) + + +def choose_num_threads(logical_cores: int | None) -> int: + if not logical_cores: + return _MIN_THREADS + return max(_MIN_THREADS, logical_cores - _SPARE_CORES) diff --git a/src/vox/adapters/diarization_models.py b/src/vox/adapters/diarization_models.py new file mode 100644 index 0000000..a8243f7 --- /dev/null +++ b/src/vox/adapters/diarization_models.py @@ -0,0 +1,25 @@ +from vox.models.exceptions import DiarizationError + +SEGMENTATION_REPO = "csukuangfj/sherpa-onnx-pyannote-segmentation-3-0" +SEGMENTATION_FILE = "model.onnx" +EMBEDDING_REPO = "csukuangfj/speaker-embedding-models" +EMBEDDING_FILE = "wespeaker_en_voxceleb_resnet34_LM.onnx" + + +def ensure_model(repo: str, filename: str) -> str: + from huggingface_hub import hf_hub_download + + try: + return hf_hub_download(repo_id=repo, filename=filename) + except Exception as e: + raise DiarizationError(f"cannot download {filename} from {repo}: {e}") from e + + +def import_sherpa(): + try: + import sherpa_onnx + except ImportError as e: + raise DiarizationError( + "sherpa-onnx not installed. Run: uv add sherpa-onnx sherpa-onnx-core" + ) from e + return sherpa_onnx diff --git a/src/vox/adapters/disk_file_writer.py b/src/vox/adapters/disk_file_writer.py index 860a748..65dda25 100644 --- a/src/vox/adapters/disk_file_writer.py +++ b/src/vox/adapters/disk_file_writer.py @@ -9,7 +9,7 @@ def write_srt(self, result: TranscriptionResult, path: Path) -> None: path.write_text(_format_srt(result), encoding="utf-8") def write_txt(self, result: TranscriptionResult, path: Path) -> None: - path.write_text(result.text, encoding="utf-8") + path.write_text(_format_txt(result), encoding="utf-8") def write_json(self, result: TranscriptionResult, path: Path) -> None: payload = _to_dict(result) @@ -30,7 +30,32 @@ def _format_srt(result: TranscriptionResult) -> str: def _format_srt_block(index: int, segment) -> str: start = _seconds_to_srt_timecode(segment.start) end = _seconds_to_srt_timecode(segment.end) - return f"{index}\n{start} --> {end}\n{segment.text}" + return f"{index}\n{start} --> {end}\n{_labelled_text(segment)}" + + +def _labelled_text(segment) -> str: + if segment.speaker is None: + return segment.text + return f"[{segment.speaker}] {segment.text}" + + +def _format_txt(result: TranscriptionResult) -> str: + if not any(s.speaker for s in result.segments): + return result.text + return "\n\n".join( + f"{speaker}: {text}" for speaker, text in _group_by_speaker(result.segments) + ) + + +def _group_by_speaker(segments) -> list[tuple[str, str]]: + blocks: list[tuple[str, list[str]]] = [] + for segment in segments: + speaker = segment.speaker or "UNKNOWN" + if blocks and blocks[-1][0] == speaker: + blocks[-1][1].append(segment.text.strip()) + else: + blocks.append((speaker, [segment.text.strip()])) + return [(speaker, " ".join(parts)) for speaker, parts in blocks] def _seconds_to_srt_timecode(seconds: float) -> str: @@ -54,6 +79,7 @@ def _segment_to_dict(segment) -> dict: "start": segment.start, "end": segment.end, "text": segment.text, + "speaker": segment.speaker, } diff --git a/src/vox/adapters/json_voice_print_store.py b/src/vox/adapters/json_voice_print_store.py new file mode 100644 index 0000000..b388dbc --- /dev/null +++ b/src/vox/adapters/json_voice_print_store.py @@ -0,0 +1,24 @@ +import json +from pathlib import Path + +from vox.models.voice_print import VoicePrint + + +class JsonVoicePrintStore: + def __init__(self, path: Path | None = None): + self._path = path or Path.home() / ".vox" / "voiceprints.json" + + def save(self, voice_print: VoicePrint) -> None: + prints = {p.name: list(p.embedding) for p in self.load_all()} + prints[voice_print.name] = list(voice_print.embedding) + self._path.parent.mkdir(parents=True, exist_ok=True) + self._path.write_text(json.dumps(prints, indent=2), encoding="utf-8") + + def load_all(self) -> tuple[VoicePrint, ...]: + if not self._path.exists(): + return () + raw = json.loads(self._path.read_text(encoding="utf-8")) + return tuple( + VoicePrint(name=name, embedding=tuple(embedding)) + for name, embedding in raw.items() + ) diff --git a/src/vox/adapters/sherpa_diarizer.py b/src/vox/adapters/sherpa_diarizer.py new file mode 100644 index 0000000..961f28d --- /dev/null +++ b/src/vox/adapters/sherpa_diarizer.py @@ -0,0 +1,75 @@ +from pathlib import Path + +from vox.adapters.audio_decoding import decode_to_mono_16k +from vox.adapters.cpu_threads import default_num_threads +from vox.adapters.diarization_models import ( + EMBEDDING_FILE, + EMBEDDING_REPO, + SEGMENTATION_FILE, + SEGMENTATION_REPO, + ensure_model, + import_sherpa, +) +from vox.models.exceptions import DiarizationError +from vox.models.speaker_turn import SpeakerTurn + +_AUTO_CLUSTERS = -1 + + +class SherpaDiarizer: + def __init__( + self, + num_threads: int | None = None, + clustering_threshold: float = 0.5, + ): + self._num_threads = num_threads or default_num_threads() + self._clustering_threshold = clustering_threshold + + def diarize( + self, + audio_path: Path, + num_speakers: int | None, + ) -> tuple[SpeakerTurn, ...]: + samples = decode_to_mono_16k(audio_path) + pipeline = self._build_pipeline(num_speakers) + try: + result = pipeline.process(samples) + except Exception as e: + raise DiarizationError(str(e)) from e + return _to_turns(result.sort_by_start_time()) + + def _build_pipeline(self, num_speakers: int | None): + sherpa_onnx = import_sherpa() + config = sherpa_onnx.OfflineSpeakerDiarizationConfig( + segmentation=sherpa_onnx.OfflineSpeakerSegmentationModelConfig( + pyannote=sherpa_onnx.OfflineSpeakerSegmentationPyannoteModelConfig( + model=ensure_model(SEGMENTATION_REPO, SEGMENTATION_FILE), + ), + num_threads=self._num_threads, + ), + embedding=sherpa_onnx.SpeakerEmbeddingExtractorConfig( + model=ensure_model(EMBEDDING_REPO, EMBEDDING_FILE), + num_threads=self._num_threads, + ), + clustering=sherpa_onnx.FastClusteringConfig( + num_clusters=num_speakers or _AUTO_CLUSTERS, + threshold=self._clustering_threshold, + ), + ) + return sherpa_onnx.OfflineSpeakerDiarization(config) + + +def _to_turns(raw_segments) -> tuple[SpeakerTurn, ...]: + return tuple( + SpeakerTurn( + start=s.start, + end=s.end, + speaker=_speaker_label(s.speaker), + ) + for s in raw_segments + if s.end > s.start + ) + + +def _speaker_label(speaker: int) -> str: + return f"SPEAKER_{int(speaker):02d}" diff --git a/src/vox/adapters/sherpa_voice_print_extractor.py b/src/vox/adapters/sherpa_voice_print_extractor.py new file mode 100644 index 0000000..4cc05bc --- /dev/null +++ b/src/vox/adapters/sherpa_voice_print_extractor.py @@ -0,0 +1,43 @@ +from pathlib import Path + +from vox.adapters.audio_decoding import SAMPLE_RATE, decode_to_mono_16k, slice_samples +from vox.adapters.cpu_threads import default_num_threads +from vox.adapters.diarization_models import ( + EMBEDDING_FILE, + EMBEDDING_REPO, + ensure_model, + import_sherpa, +) +from vox.models.exceptions import DiarizationError + + +class SherpaVoicePrintExtractor: + def __init__(self, num_threads: int | None = None): + self._num_threads = num_threads or default_num_threads() + self._extractor = None + + def extract( + self, + audio_path: Path, + start: float, + end: float, + ) -> tuple[float, ...]: + samples = slice_samples(decode_to_mono_16k(audio_path), start, end) + if len(samples) == 0: + raise DiarizationError(f"no audio between {start}s and {end}s") + extractor = self._ensure_extractor() + stream = extractor.create_stream() + stream.accept_waveform(SAMPLE_RATE, samples) + stream.input_finished() + return tuple(extractor.compute(stream)) + + def _ensure_extractor(self): + if self._extractor is None: + sherpa_onnx = import_sherpa() + self._extractor = sherpa_onnx.SpeakerEmbeddingExtractor( + sherpa_onnx.SpeakerEmbeddingExtractorConfig( + model=ensure_model(EMBEDDING_REPO, EMBEDDING_FILE), + num_threads=self._num_threads, + ) + ) + return self._extractor diff --git a/src/vox/adapters/system_dep_checker.py b/src/vox/adapters/system_dep_checker.py index 4a908c1..e86b60a 100644 --- a/src/vox/adapters/system_dep_checker.py +++ b/src/vox/adapters/system_dep_checker.py @@ -10,6 +10,7 @@ def check_all(self) -> list[HealthStatus]: _check_python_module("yt_dlp", "yt-dlp"), _check_binary("ffmpeg"), _check_python_module("mlx_whisper", "mlx-whisper"), + _check_python_module("sherpa_onnx", "sherpa-onnx"), ] diff --git a/src/vox/models/exceptions.py b/src/vox/models/exceptions.py index 1d09a36..df90e27 100644 --- a/src/vox/models/exceptions.py +++ b/src/vox/models/exceptions.py @@ -22,6 +22,10 @@ class TranscriptionError(VoxError): pass +class DiarizationError(VoxError): + pass + + class ConfigError(VoxError): pass diff --git a/src/vox/models/segment.py b/src/vox/models/segment.py index 9b65bd3..0286e60 100644 --- a/src/vox/models/segment.py +++ b/src/vox/models/segment.py @@ -8,6 +8,7 @@ class Segment: start: float end: float text: str + speaker: str | None = None def __post_init__(self): if self.start < 0: diff --git a/src/vox/models/speaker_assignment.py b/src/vox/models/speaker_assignment.py new file mode 100644 index 0000000..484647d --- /dev/null +++ b/src/vox/models/speaker_assignment.py @@ -0,0 +1,31 @@ +from dataclasses import replace + +from vox.models.segment import Segment +from vox.models.speaker_turn import SpeakerTurn + + +def assign_speakers( + segments: tuple[Segment, ...], + turns: tuple[SpeakerTurn, ...], +) -> tuple[Segment, ...]: + if not turns: + return segments + return tuple(replace(s, speaker=_dominant_speaker(s, turns)) for s in segments) + + +def _dominant_speaker( + segment: Segment, + turns: tuple[SpeakerTurn, ...], +) -> str | None: + totals: dict[str, float] = {} + for turn in turns: + overlap = _overlap_seconds(segment, turn) + if overlap > 0: + totals[turn.speaker] = totals.get(turn.speaker, 0.0) + overlap + if not totals: + return None + return max(totals, key=lambda speaker: totals[speaker]) + + +def _overlap_seconds(segment: Segment, turn: SpeakerTurn) -> float: + return min(segment.end, turn.end) - max(segment.start, turn.start) diff --git a/src/vox/models/speaker_renaming.py b/src/vox/models/speaker_renaming.py new file mode 100644 index 0000000..357ad9b --- /dev/null +++ b/src/vox/models/speaker_renaming.py @@ -0,0 +1,18 @@ +from dataclasses import replace + +from vox.models.segment import Segment + + +def rename_speakers( + segments: tuple[Segment, ...], + mapping: dict[str, str], +) -> tuple[Segment, ...]: + if not mapping: + return segments + return tuple(replace(s, speaker=_renamed(s.speaker, mapping)) for s in segments) + + +def _renamed(speaker: str | None, mapping: dict[str, str]) -> str | None: + if speaker is None: + return None + return mapping.get(speaker, speaker) diff --git a/src/vox/models/speaker_turn.py b/src/vox/models/speaker_turn.py new file mode 100644 index 0000000..e6a60c4 --- /dev/null +++ b/src/vox/models/speaker_turn.py @@ -0,0 +1,18 @@ +from dataclasses import dataclass + +from vox.models.exceptions import ValidationError + + +@dataclass(frozen=True) +class SpeakerTurn: + start: float + end: float + speaker: str + + def __post_init__(self): + if self.start < 0: + raise ValidationError("start must be >= 0") + if self.end <= self.start: + raise ValidationError("end must be > start") + if not self.speaker.strip(): + raise ValidationError("speaker must not be blank") diff --git a/src/vox/models/turn_selection.py b/src/vox/models/turn_selection.py new file mode 100644 index 0000000..2f850e9 --- /dev/null +++ b/src/vox/models/turn_selection.py @@ -0,0 +1,16 @@ +from vox.models.speaker_turn import SpeakerTurn + + +def longest_turn_per_speaker( + turns: tuple[SpeakerTurn, ...], +) -> dict[str, SpeakerTurn]: + longest: dict[str, SpeakerTurn] = {} + for turn in turns: + current = longest.get(turn.speaker) + if current is None or _duration(turn) > _duration(current): + longest[turn.speaker] = turn + return longest + + +def _duration(turn: SpeakerTurn) -> float: + return turn.end - turn.start diff --git a/src/vox/models/voice_matching.py b/src/vox/models/voice_matching.py new file mode 100644 index 0000000..03b7839 --- /dev/null +++ b/src/vox/models/voice_matching.py @@ -0,0 +1,43 @@ +import math + +from vox.models.voice_print import VoicePrint + +Embedding = tuple[float, ...] + + +def cosine_similarity(left: Embedding, right: Embedding) -> float: + norm = _norm(left) * _norm(right) + if norm == 0: + return 0.0 + return sum(a * b for a, b in zip(left, right, strict=True)) / norm + + +def match_speakers( + label_embeddings: dict[str, Embedding], + prints: tuple[VoicePrint, ...], + threshold: float, +) -> dict[str, str]: + if not prints: + return {} + mapping = {} + for label, embedding in label_embeddings.items(): + name = _best_match(embedding, prints, threshold) + if name is not None: + mapping[label] = name + return mapping + + +def _best_match( + embedding: Embedding, + prints: tuple[VoicePrint, ...], + threshold: float, +) -> str | None: + scored = [(cosine_similarity(embedding, p.embedding), p.name) for p in prints] + best_score, best_name = max(scored) + if best_score < threshold: + return None + return best_name + + +def _norm(vector: Embedding) -> float: + return math.sqrt(sum(v * v for v in vector)) diff --git a/src/vox/models/voice_print.py b/src/vox/models/voice_print.py new file mode 100644 index 0000000..e8d0595 --- /dev/null +++ b/src/vox/models/voice_print.py @@ -0,0 +1,15 @@ +from dataclasses import dataclass + +from vox.models.exceptions import ValidationError + + +@dataclass(frozen=True) +class VoicePrint: + name: str + embedding: tuple[float, ...] + + def __post_init__(self): + if not self.name.strip(): + raise ValidationError("name must not be blank") + if not self.embedding: + raise ValidationError("embedding must not be empty") diff --git a/src/vox/ports/diarizer.py b/src/vox/ports/diarizer.py new file mode 100644 index 0000000..701823d --- /dev/null +++ b/src/vox/ports/diarizer.py @@ -0,0 +1,12 @@ +from pathlib import Path +from typing import Protocol + +from vox.models.speaker_turn import SpeakerTurn + + +class Diarizer(Protocol): + def diarize( + self, + audio_path: Path, + num_speakers: int | None, + ) -> tuple[SpeakerTurn, ...]: ... diff --git a/src/vox/ports/voice_print_extractor.py b/src/vox/ports/voice_print_extractor.py new file mode 100644 index 0000000..e9b1ff5 --- /dev/null +++ b/src/vox/ports/voice_print_extractor.py @@ -0,0 +1,11 @@ +from pathlib import Path +from typing import Protocol + + +class VoicePrintExtractor(Protocol): + def extract( + self, + audio_path: Path, + start: float, + end: float, + ) -> tuple[float, ...]: ... diff --git a/src/vox/ports/voice_print_store.py b/src/vox/ports/voice_print_store.py new file mode 100644 index 0000000..3a2be15 --- /dev/null +++ b/src/vox/ports/voice_print_store.py @@ -0,0 +1,9 @@ +from typing import Protocol + +from vox.models.voice_print import VoicePrint + + +class VoicePrintStore(Protocol): + def save(self, voice_print: VoicePrint) -> None: ... + + def load_all(self) -> tuple[VoicePrint, ...]: ... diff --git a/src/vox/schemas/transcribe.json b/src/vox/schemas/transcribe.json index 2591463..1a64891 100644 --- a/src/vox/schemas/transcribe.json +++ b/src/vox/schemas/transcribe.json @@ -57,6 +57,20 @@ "json": { "type": "string", "description": "Raw JSON payload with keys: input (source), language, model. Overrides CLI arguments." + }, + "diarize": { + "type": "boolean", + "default": false, + "description": "Identify speakers: each segment is labelled with SPEAKER_00, SPEAKER_01, etc." + }, + "speakers": { + "type": "integer", + "description": "Known speaker count. Auto-detected when omitted, but passing the exact count is markedly more reliable." + }, + "identify": { + "type": "boolean", + "default": false, + "description": "Replace SPEAKER_00 labels with real names, matched against voices registered via 'vox speakers add'. Requires diarize." } }, "output": { diff --git a/src/vox/use_cases/enroll_speaker.py b/src/vox/use_cases/enroll_speaker.py new file mode 100644 index 0000000..52a5653 --- /dev/null +++ b/src/vox/use_cases/enroll_speaker.py @@ -0,0 +1,53 @@ +from dataclasses import dataclass +from pathlib import Path + +from vox.models.exceptions import ValidationError +from vox.models.voice_print import VoicePrint +from vox.ports.voice_print_extractor import VoicePrintExtractor +from vox.ports.voice_print_store import VoicePrintStore + +MIN_ENROLL_SECONDS = 2.0 + + +@dataclass(frozen=True) +class EnrollSpeakerRequest: + name: str + audio_path: str + start: float + end: float + + +@dataclass(frozen=True) +class EnrollSpeakerResponse: + name: str + seconds_used: float + + +class EnrollSpeakerUseCase: + def __init__( + self, + extractor: VoicePrintExtractor, + store: VoicePrintStore, + ): + self._extractor = extractor + self._store = store + + def execute(self, request: EnrollSpeakerRequest) -> EnrollSpeakerResponse: + _validate_span(request.start, request.end) + embedding = self._extractor.extract( + Path(request.audio_path), request.start, request.end + ) + self._store.save(VoicePrint(name=request.name, embedding=embedding)) + return EnrollSpeakerResponse( + name=request.name, + seconds_used=request.end - request.start, + ) + + +def _validate_span(start: float, end: float) -> None: + if end <= start: + raise ValidationError("end must be > start") + if end - start < MIN_ENROLL_SECONDS: + raise ValidationError( + f"need at least {MIN_ENROLL_SECONDS} seconds of speech to enroll a voice" + ) diff --git a/src/vox/use_cases/identify_speakers.py b/src/vox/use_cases/identify_speakers.py new file mode 100644 index 0000000..c568098 --- /dev/null +++ b/src/vox/use_cases/identify_speakers.py @@ -0,0 +1,42 @@ +from pathlib import Path + +from vox.models.speaker_turn import SpeakerTurn +from vox.models.turn_selection import longest_turn_per_speaker +from vox.models.voice_matching import match_speakers +from vox.ports.voice_print_extractor import VoicePrintExtractor +from vox.ports.voice_print_store import VoicePrintStore + +DEFAULT_MATCH_THRESHOLD = 0.8 + + +class IdentifySpeakersUseCase: + def __init__( + self, + extractor: VoicePrintExtractor, + store: VoicePrintStore, + threshold: float = DEFAULT_MATCH_THRESHOLD, + ): + self._extractor = extractor + self._store = store + self._threshold = threshold + + def execute( + self, + audio_path: Path, + turns: tuple[SpeakerTurn, ...], + ) -> dict[str, str]: + prints = self._store.load_all() + if not prints: + return {} + embeddings = self._embed_each_speaker(audio_path, turns) + return match_speakers(embeddings, prints, self._threshold) + + def _embed_each_speaker( + self, + audio_path: Path, + turns: tuple[SpeakerTurn, ...], + ) -> dict[str, tuple[float, ...]]: + return { + label: self._extractor.extract(audio_path, turn.start, turn.end) + for label, turn in longest_turn_per_speaker(turns).items() + } diff --git a/src/vox/use_cases/transcribe.py b/src/vox/use_cases/transcribe.py index e844bc9..49a320a 100644 --- a/src/vox/use_cases/transcribe.py +++ b/src/vox/use_cases/transcribe.py @@ -1,18 +1,22 @@ from __future__ import annotations import time -from dataclasses import dataclass +from dataclasses import dataclass, replace from pathlib import Path from vox.models.audio_config import AudioConfig from vox.models.exceptions import ValidationError from vox.models.language import Language +from vox.models.speaker_assignment import assign_speakers +from vox.models.speaker_renaming import rename_speakers from vox.models.transcription_input import TranscriptionInput from vox.ports.audio_cleaner import AudioCleaner +from vox.ports.diarizer import Diarizer from vox.ports.downloader import Downloader from vox.ports.file_writer import FileWriter from vox.ports.progress_reporter import ProgressReporter from vox.ports.transcriber import Transcriber +from vox.use_cases.identify_speakers import IdentifySpeakersUseCase @dataclass(frozen=True) @@ -26,6 +30,9 @@ class TranscribeRequest: no_download: bool = False dry_run: bool = False output_stem: str = "" + diarize: bool = False + identify: bool = False + num_speakers: int | None = None @dataclass(frozen=True) @@ -46,12 +53,16 @@ def __init__( transcriber: Transcriber, file_writer: FileWriter, progress: ProgressReporter, + diarizer: Diarizer | None = None, + speaker_identifier: IdentifySpeakersUseCase | None = None, ): self._downloader = downloader self._audio_cleaner = audio_cleaner self._transcriber = transcriber self._file_writer = file_writer self._progress = progress + self._diarizer = diarizer + self._speaker_identifier = speaker_identifier def execute(self, request: TranscribeRequest) -> TranscribeResponse: self._progress.start("Validating input") @@ -64,10 +75,10 @@ def execute(self, request: TranscribeRequest) -> TranscribeResponse: return _dry_run_response(request, parsed_input, output_dir) audio_path = self._resolve_audio(parsed_input, output_dir) - wav_path = self._maybe_clean(audio_path, request.no_clean, output_dir) - result = self._transcribe( - wav_path or audio_path, request.model, language, request - ) + wav_path = self._maybe_clean(audio_path, request, output_dir) + transcribed_path = wav_path or audio_path + result = self._transcribe(transcribed_path, request.model, language, request) + result = self._maybe_diarize(result, transcribed_path, request) paths = self._write_outputs( result, output_dir, parsed_input, request.output_stem ) @@ -82,14 +93,13 @@ def _resolve_audio( self._progress.update("Downloading") return self._downloader.download(parsed_input, output_dir) - def _maybe_clean( - self, audio_path: Path, no_clean: bool, output_dir: Path - ) -> Path | None: - if no_clean: + def _maybe_clean(self, audio_path: Path, request, output_dir: Path) -> Path | None: + if request.no_clean: return None self._progress.update("Cleaning audio") clean_path = output_dir / f"{audio_path.stem}_clean.wav" - return self._audio_cleaner.clean(audio_path, AudioConfig.default(), clean_path) + config = _audio_config_for(request.diarize) + return self._audio_cleaner.clean(audio_path, config, clean_path) def _transcribe(self, audio_path, model, language, request): self._progress.update("Transcribing") @@ -98,6 +108,24 @@ def _transcribe(self, audio_path, model, language, request): audio_path, model, lang_code, request.word_timestamps ) + def _maybe_diarize(self, result, audio_path: Path, request): + if not request.diarize: + return result + if self._diarizer is None: + raise ValidationError("Diarization requested but no diarizer configured") + self._progress.update("Identifying speakers") + turns = self._diarizer.diarize(audio_path, request.num_speakers) + segments = assign_speakers(result.segments, turns) + segments = self._maybe_identify(segments, turns, audio_path, request) + return replace(result, segments=segments) + + def _maybe_identify(self, segments, turns, audio_path: Path, request): + if not request.identify or self._speaker_identifier is None: + return segments + self._progress.update("Matching known voices") + mapping = self._speaker_identifier.execute(audio_path, turns) + return rename_speakers(segments, mapping) + def _write_outputs(self, result, output_dir, parsed_input, output_stem): self._progress.update("Writing outputs") stem = output_stem or _derive_output_stem(parsed_input) @@ -110,6 +138,12 @@ def _write_outputs(self, result, output_dir, parsed_input, output_stem): return srt_path, txt_path, json_path +def _audio_config_for(diarize: bool) -> AudioConfig: + if not diarize: + return AudioConfig.default() + return AudioConfig(remove_silence=False, denoise=False) + + def _reject_no_download_with_url( no_download: bool, parsed_input: TranscriptionInput ) -> None: @@ -151,6 +185,8 @@ def _build_execution_plan( steps.append("Clean audio via ffmpeg") lang = request.language steps.append(f"Transcribe with model '{request.model}', language '{lang}'") + if request.diarize: + steps.append("Identify speakers") steps.append("Write outputs: .srt, .txt, .json") numbered = [f"{i}. {s}" for i, s in enumerate(steps, 1)] return "Execution plan:\n" + "\n".join(numbered) diff --git a/tests/fakes/fake_diarizer.py b/tests/fakes/fake_diarizer.py new file mode 100644 index 0000000..ecc7e21 --- /dev/null +++ b/tests/fakes/fake_diarizer.py @@ -0,0 +1,21 @@ +from pathlib import Path + +from vox.models.speaker_turn import SpeakerTurn + + +class FakeDiarizer: + def __init__(self, turns: tuple[SpeakerTurn, ...] | None = None): + self._turns = turns if turns is not None else _default_turns() + self.diarize_called_with: tuple | None = None + + def diarize( + self, + audio_path: Path, + num_speakers: int | None, + ) -> tuple[SpeakerTurn, ...]: + self.diarize_called_with = (audio_path, num_speakers) + return self._turns + + +def _default_turns() -> tuple[SpeakerTurn, ...]: + return (SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"),) diff --git a/tests/fakes/fake_speaker_identifier.py b/tests/fakes/fake_speaker_identifier.py new file mode 100644 index 0000000..64db097 --- /dev/null +++ b/tests/fakes/fake_speaker_identifier.py @@ -0,0 +1,17 @@ +from pathlib import Path + +from vox.models.speaker_turn import SpeakerTurn + + +class FakeSpeakerIdentifier: + def __init__(self, mapping: dict[str, str] | None = None): + self.mapping = mapping or {} + self.called_with: tuple[Path, tuple[SpeakerTurn, ...]] | None = None + + def execute( + self, + audio_path: Path, + turns: tuple[SpeakerTurn, ...], + ) -> dict[str, str]: + self.called_with = (audio_path, turns) + return self.mapping diff --git a/tests/fakes/fake_voice_print_extractor.py b/tests/fakes/fake_voice_print_extractor.py new file mode 100644 index 0000000..8727893 --- /dev/null +++ b/tests/fakes/fake_voice_print_extractor.py @@ -0,0 +1,18 @@ +from pathlib import Path + + +class FakeVoicePrintExtractor: + def __init__( + self, embeddings: dict[tuple[float, float], tuple[float, ...]] | None = None + ): + self._embeddings = embeddings or {} + self.extract_calls: list[tuple[Path, float, float]] = [] + + def extract( + self, + audio_path: Path, + start: float, + end: float, + ) -> tuple[float, ...]: + self.extract_calls.append((audio_path, start, end)) + return self._embeddings.get((start, end), (1.0, 0.0)) diff --git a/tests/fakes/fake_voice_print_store.py b/tests/fakes/fake_voice_print_store.py new file mode 100644 index 0000000..fb2ac02 --- /dev/null +++ b/tests/fakes/fake_voice_print_store.py @@ -0,0 +1,15 @@ +from vox.models.voice_print import VoicePrint + + +class FakeVoicePrintStore: + def __init__(self, prints: tuple[VoicePrint, ...] = ()): + self._prints = list(prints) + self.saved: list[VoicePrint] = [] + + def save(self, voice_print: VoicePrint) -> None: + self.saved.append(voice_print) + self._prints = [p for p in self._prints if p.name != voice_print.name] + self._prints.append(voice_print) + + def load_all(self) -> tuple[VoicePrint, ...]: + return tuple(self._prints) diff --git a/tests/unit/adapters/test_cpu_threads.py b/tests/unit/adapters/test_cpu_threads.py new file mode 100644 index 0000000..78513bf --- /dev/null +++ b/tests/unit/adapters/test_cpu_threads.py @@ -0,0 +1,24 @@ +from vox.adapters.cpu_threads import choose_num_threads + + +class TestChooseNumThreads: + def test_choose_when_ten_cores_then_uses_eight(self): + assert choose_num_threads(logical_cores=10) == 8 + + def test_choose_when_many_cores_then_leaves_two_spare(self): + assert choose_num_threads(logical_cores=16) == 14 + + def test_choose_when_four_cores_then_two(self): + assert choose_num_threads(logical_cores=4) == 2 + + def test_choose_when_three_cores_then_stays_above_one(self): + assert choose_num_threads(logical_cores=3) == 2 + + def test_choose_when_two_cores_then_stays_above_one(self): + assert choose_num_threads(logical_cores=2) == 2 + + def test_choose_when_single_core_then_stays_above_one(self): + assert choose_num_threads(logical_cores=1) == 2 + + def test_choose_when_unknown_then_stays_above_one(self): + assert choose_num_threads(logical_cores=None) == 2 diff --git a/tests/unit/adapters/test_disk_file_writer.py b/tests/unit/adapters/test_disk_file_writer.py index 909c01b..8d1c6d1 100644 --- a/tests/unit/adapters/test_disk_file_writer.py +++ b/tests/unit/adapters/test_disk_file_writer.py @@ -91,3 +91,81 @@ def test_write_json_when_called_then_valid_json(self, tmp_path): assert parsed["segments"][0]["end"] == 2.5 assert parsed["segments"][0]["text"] == "Hello world" assert parsed["words"] is None + + +class TestFormatSrtWithSpeakers: + def test_format_srt_when_segment_has_speaker_then_prefixes_label(self): + result = _make_result( + segments=( + Segment(start=0.0, end=2.5, text="Hello world", speaker="SPEAKER_00"), + ), + ) + + srt = _format_srt(result) + + assert srt == ("1\n00:00:00,000 --> 00:00:02,500\n[SPEAKER_00] Hello world\n") + + def test_format_srt_when_segment_has_no_speaker_then_text_unchanged(self): + result = _make_result( + segments=(Segment(start=0.0, end=2.5, text="Hello world"),), + ) + + srt = _format_srt(result) + + assert "[" not in srt + + +class TestWriteJsonWithSpeakers: + def test_write_json_when_speaker_present_then_included(self, tmp_path): + result = _make_result( + segments=( + Segment(start=0.0, end=2.5, text="Hello world", speaker="SPEAKER_01"), + ), + ) + path = tmp_path / "output.json" + + DiskFileWriter().write_json(result, path) + + parsed = json.loads(path.read_text(encoding="utf-8")) + assert parsed["segments"][0]["speaker"] == "SPEAKER_01" + + def test_write_json_when_no_speaker_then_field_is_none(self, tmp_path): + result = _make_result( + segments=(Segment(start=0.0, end=2.5, text="Hello world"),), + ) + path = tmp_path / "output.json" + + DiskFileWriter().write_json(result, path) + + parsed = json.loads(path.read_text(encoding="utf-8")) + assert parsed["segments"][0]["speaker"] is None + + +class TestWriteTxtWithSpeakers: + def test_write_txt_when_speakers_present_then_groups_by_speaker(self, tmp_path): + result = _make_result( + segments=( + Segment(start=0.0, end=1.0, text="Bonjour", speaker="SPEAKER_00"), + Segment(start=1.0, end=2.0, text="tout le monde", speaker="SPEAKER_00"), + Segment(start=2.0, end=3.0, text="Salut", speaker="SPEAKER_01"), + ), + text="Bonjour tout le monde Salut", + ) + path = tmp_path / "output.txt" + + DiskFileWriter().write_txt(result, path) + + assert path.read_text(encoding="utf-8") == ( + "SPEAKER_00: Bonjour tout le monde\n\nSPEAKER_01: Salut" + ) + + def test_write_txt_when_no_speakers_then_plain_text_unchanged(self, tmp_path): + result = _make_result( + segments=(Segment(start=0.0, end=1.0, text="Hi"),), + text="Hello world", + ) + path = tmp_path / "output.txt" + + DiskFileWriter().write_txt(result, path) + + assert path.read_text(encoding="utf-8") == "Hello world" diff --git a/tests/unit/adapters/test_json_voice_print_store.py b/tests/unit/adapters/test_json_voice_print_store.py new file mode 100644 index 0000000..0884679 --- /dev/null +++ b/tests/unit/adapters/test_json_voice_print_store.py @@ -0,0 +1,37 @@ +from vox.adapters.json_voice_print_store import JsonVoicePrintStore +from vox.models.voice_print import VoicePrint + + +class TestJsonVoicePrintStore: + def test_load_all_when_no_file_then_empty(self, tmp_path): + store = JsonVoicePrintStore(tmp_path / "prints.json") + + assert store.load_all() == () + + def test_save_then_load_all_returns_print(self, tmp_path): + store = JsonVoicePrintStore(tmp_path / "prints.json") + + store.save(VoicePrint(name="Coco", embedding=(0.1, 0.2))) + + loaded = store.load_all() + assert len(loaded) == 1 + assert loaded[0].name == "Coco" + assert loaded[0].embedding == (0.1, 0.2) + + def test_save_when_same_name_twice_then_overwrites(self, tmp_path): + store = JsonVoicePrintStore(tmp_path / "prints.json") + + store.save(VoicePrint(name="Coco", embedding=(0.1,))) + store.save(VoicePrint(name="Coco", embedding=(0.9,))) + + loaded = store.load_all() + assert len(loaded) == 1 + assert loaded[0].embedding == (0.9,) + + def test_save_when_different_names_then_both_kept(self, tmp_path): + store = JsonVoicePrintStore(tmp_path / "prints.json") + + store.save(VoicePrint(name="Coco", embedding=(0.1,))) + store.save(VoicePrint(name="Leslie", embedding=(0.2,))) + + assert {p.name for p in store.load_all()} == {"Coco", "Leslie"} diff --git a/tests/unit/adapters/test_sherpa_diarizer.py b/tests/unit/adapters/test_sherpa_diarizer.py new file mode 100644 index 0000000..88e64cc --- /dev/null +++ b/tests/unit/adapters/test_sherpa_diarizer.py @@ -0,0 +1,38 @@ +from types import SimpleNamespace + +from vox.adapters.sherpa_diarizer import _speaker_label, _to_turns + + +def _raw_segment(start, end, speaker): + return SimpleNamespace(start=start, end=end, speaker=speaker) + + +class TestSpeakerLabel: + def test_label_when_zero_then_padded(self): + assert _speaker_label(0) == "SPEAKER_00" + + def test_label_when_two_digits_then_not_padded(self): + assert _speaker_label(12) == "SPEAKER_12" + + +class TestToTurns: + def test_to_turns_when_segments_then_maps_fields(self): + raw = [_raw_segment(0.0, 2.5, 0), _raw_segment(2.5, 4.0, 1)] + + turns = _to_turns(raw) + + assert len(turns) == 2 + assert turns[0].start == 0.0 + assert turns[0].end == 2.5 + assert turns[0].speaker == "SPEAKER_00" + assert turns[1].speaker == "SPEAKER_01" + + def test_to_turns_when_zero_length_segment_then_skipped(self): + raw = [_raw_segment(1.0, 1.0, 0), _raw_segment(1.0, 2.0, 0)] + + turns = _to_turns(raw) + + assert len(turns) == 1 + + def test_to_turns_when_empty_then_returns_empty_tuple(self): + assert _to_turns([]) == () diff --git a/tests/unit/models/test_segment.py b/tests/unit/models/test_segment.py index df67080..69f14c8 100644 --- a/tests/unit/models/test_segment.py +++ b/tests/unit/models/test_segment.py @@ -23,3 +23,13 @@ def test_create_when_end_before_start_then_raises(self): def test_create_when_end_equals_start_then_raises(self): with pytest.raises(ValidationError, match="end"): Segment(start=1.0, end=1.0, text="hello") + + def test_create_when_no_speaker_then_speaker_is_none(self): + segment = Segment(start=0.0, end=1.5, text="hello") + + assert segment.speaker is None + + def test_create_when_speaker_given_then_speaker_set(self): + segment = Segment(start=0.0, end=1.5, text="hello", speaker="SPEAKER_01") + + assert segment.speaker == "SPEAKER_01" diff --git a/tests/unit/models/test_speaker_assignment.py b/tests/unit/models/test_speaker_assignment.py new file mode 100644 index 0000000..359f962 --- /dev/null +++ b/tests/unit/models/test_speaker_assignment.py @@ -0,0 +1,72 @@ +from vox.models.segment import Segment +from vox.models.speaker_assignment import assign_speakers +from vox.models.speaker_turn import SpeakerTurn + + +class TestAssignSpeakers: + def test_assign_when_no_turns_then_segments_keep_no_speaker(self): + segments = (Segment(start=0.0, end=2.0, text="hello"),) + + result = assign_speakers(segments, ()) + + assert result[0].speaker is None + + def test_assign_when_turn_covers_segment_then_speaker_set(self): + segments = (Segment(start=1.0, end=2.0, text="hello"),) + turns = (SpeakerTurn(start=0.0, end=5.0, speaker="SPEAKER_00"),) + + result = assign_speakers(segments, turns) + + assert result[0].speaker == "SPEAKER_00" + + def test_assign_when_segment_straddles_two_turns_then_takes_dominant(self): + segments = (Segment(start=8.0, end=14.0, text="hello"),) + turns = ( + SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=20.0, speaker="SPEAKER_01"), + ) + + result = assign_speakers(segments, turns) + + assert result[0].speaker == "SPEAKER_01" + + def test_assign_when_segment_outside_all_turns_then_speaker_is_none(self): + segments = (Segment(start=30.0, end=31.0, text="hello"),) + turns = (SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"),) + + result = assign_speakers(segments, turns) + + assert result[0].speaker is None + + def test_assign_when_many_segments_then_each_gets_own_speaker(self): + segments = ( + Segment(start=0.0, end=5.0, text="first"), + Segment(start=15.0, end=18.0, text="second"), + ) + turns = ( + SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=20.0, speaker="SPEAKER_01"), + ) + + result = assign_speakers(segments, turns) + + assert [s.speaker for s in result] == ["SPEAKER_00", "SPEAKER_01"] + + def test_assign_when_turns_overlap_then_takes_longest_overlap(self): + segments = (Segment(start=0.0, end=10.0, text="hello"),) + turns = ( + SpeakerTurn(start=0.0, end=3.0, speaker="SPEAKER_00"), + SpeakerTurn(start=2.0, end=10.0, speaker="SPEAKER_01"), + ) + + result = assign_speakers(segments, turns) + + assert result[0].speaker == "SPEAKER_01" + + def test_assign_when_called_then_text_and_bounds_preserved(self): + segments = (Segment(start=1.0, end=2.0, text="hello"),) + turns = (SpeakerTurn(start=0.0, end=5.0, speaker="SPEAKER_00"),) + + result = assign_speakers(segments, turns) + + assert (result[0].start, result[0].end, result[0].text) == (1.0, 2.0, "hello") diff --git a/tests/unit/models/test_speaker_renaming.py b/tests/unit/models/test_speaker_renaming.py new file mode 100644 index 0000000..8a899ce --- /dev/null +++ b/tests/unit/models/test_speaker_renaming.py @@ -0,0 +1,39 @@ +from vox.models.segment import Segment +from vox.models.speaker_renaming import rename_speakers + + +class TestRenameSpeakers: + def test_rename_when_mapping_matches_then_speaker_replaced(self): + segments = (Segment(start=0.0, end=1.0, text="hi", speaker="SPEAKER_00"),) + + result = rename_speakers(segments, {"SPEAKER_00": "Coco"}) + + assert result[0].speaker == "Coco" + + def test_rename_when_no_mapping_then_label_kept(self): + segments = (Segment(start=0.0, end=1.0, text="hi", speaker="SPEAKER_01"),) + + result = rename_speakers(segments, {"SPEAKER_00": "Coco"}) + + assert result[0].speaker == "SPEAKER_01" + + def test_rename_when_empty_mapping_then_segments_unchanged(self): + segments = (Segment(start=0.0, end=1.0, text="hi", speaker="SPEAKER_00"),) + + result = rename_speakers(segments, {}) + + assert result == segments + + def test_rename_when_segment_has_no_speaker_then_left_alone(self): + segments = (Segment(start=0.0, end=1.0, text="hi"),) + + result = rename_speakers(segments, {"SPEAKER_00": "Coco"}) + + assert result[0].speaker is None + + def test_rename_when_called_then_text_and_bounds_preserved(self): + segments = (Segment(start=1.0, end=2.0, text="hi", speaker="SPEAKER_00"),) + + result = rename_speakers(segments, {"SPEAKER_00": "Coco"}) + + assert (result[0].start, result[0].end, result[0].text) == (1.0, 2.0, "hi") diff --git a/tests/unit/models/test_speaker_turn.py b/tests/unit/models/test_speaker_turn.py new file mode 100644 index 0000000..268c661 --- /dev/null +++ b/tests/unit/models/test_speaker_turn.py @@ -0,0 +1,29 @@ +import pytest + +from vox.models.exceptions import ValidationError +from vox.models.speaker_turn import SpeakerTurn + + +class TestSpeakerTurn: + def test_create_when_valid_then_fields_set(self): + turn = SpeakerTurn(start=0.0, end=1.5, speaker="SPEAKER_00") + + assert turn.start == 0.0 + assert turn.end == 1.5 + assert turn.speaker == "SPEAKER_00" + + def test_create_when_negative_start_then_raises(self): + with pytest.raises(ValidationError, match="start"): + SpeakerTurn(start=-1.0, end=1.0, speaker="SPEAKER_00") + + def test_create_when_end_before_start_then_raises(self): + with pytest.raises(ValidationError, match="end"): + SpeakerTurn(start=2.0, end=1.0, speaker="SPEAKER_00") + + def test_create_when_end_equals_start_then_raises(self): + with pytest.raises(ValidationError, match="end"): + SpeakerTurn(start=1.0, end=1.0, speaker="SPEAKER_00") + + def test_create_when_blank_speaker_then_raises(self): + with pytest.raises(ValidationError, match="speaker"): + SpeakerTurn(start=0.0, end=1.0, speaker=" ") diff --git a/tests/unit/models/test_turn_selection.py b/tests/unit/models/test_turn_selection.py new file mode 100644 index 0000000..33d358c --- /dev/null +++ b/tests/unit/models/test_turn_selection.py @@ -0,0 +1,29 @@ +from vox.models.speaker_turn import SpeakerTurn +from vox.models.turn_selection import longest_turn_per_speaker + + +class TestLongestTurnPerSpeaker: + def test_select_when_one_turn_each_then_returns_both(self): + turns = ( + SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), + SpeakerTurn(start=2.0, end=5.0, speaker="SPEAKER_01"), + ) + + selected = longest_turn_per_speaker(turns) + + assert set(selected) == {"SPEAKER_00", "SPEAKER_01"} + + def test_select_when_several_turns_then_keeps_longest(self): + turns = ( + SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=20.0, speaker="SPEAKER_00"), + SpeakerTurn(start=25.0, end=26.0, speaker="SPEAKER_00"), + ) + + selected = longest_turn_per_speaker(turns) + + assert selected["SPEAKER_00"].start == 10.0 + assert selected["SPEAKER_00"].end == 20.0 + + def test_select_when_empty_then_empty_dict(self): + assert longest_turn_per_speaker(()) == {} diff --git a/tests/unit/models/test_voice_matching.py b/tests/unit/models/test_voice_matching.py new file mode 100644 index 0000000..0234930 --- /dev/null +++ b/tests/unit/models/test_voice_matching.py @@ -0,0 +1,54 @@ +from vox.models.voice_matching import cosine_similarity, match_speakers +from vox.models.voice_print import VoicePrint + + +class TestCosineSimilarity: + def test_similarity_when_identical_then_one(self): + assert cosine_similarity((1.0, 0.0), (1.0, 0.0)) == 1.0 + + def test_similarity_when_orthogonal_then_zero(self): + assert cosine_similarity((1.0, 0.0), (0.0, 1.0)) == 0.0 + + def test_similarity_when_scaled_then_still_one(self): + assert cosine_similarity((1.0, 1.0), (3.0, 3.0)) == 1.0 + + def test_similarity_when_zero_vector_then_zero(self): + assert cosine_similarity((0.0, 0.0), (1.0, 1.0)) == 0.0 + + +class TestMatchSpeakers: + def test_match_when_above_threshold_then_named(self): + prints = (VoicePrint(name="Coco", embedding=(1.0, 0.0)),) + + mapping = match_speakers({"SPEAKER_00": (1.0, 0.0)}, prints, 0.5) + + assert mapping == {"SPEAKER_00": "Coco"} + + def test_match_when_below_threshold_then_absent(self): + prints = (VoicePrint(name="Coco", embedding=(1.0, 0.0)),) + + mapping = match_speakers({"SPEAKER_00": (0.0, 1.0)}, prints, 0.5) + + assert mapping == {} + + def test_match_when_several_prints_then_best_wins(self): + prints = ( + VoicePrint(name="Coco", embedding=(1.0, 0.0)), + VoicePrint(name="Leslie", embedding=(0.9, 0.4)), + ) + + mapping = match_speakers({"SPEAKER_00": (0.9, 0.4)}, prints, 0.5) + + assert mapping == {"SPEAKER_00": "Leslie"} + + def test_match_when_no_prints_then_empty(self): + assert match_speakers({"SPEAKER_00": (1.0, 0.0)}, (), 0.5) == {} + + def test_match_when_two_labels_match_same_person_then_both_named(self): + prints = (VoicePrint(name="Coco", embedding=(1.0, 0.0)),) + + mapping = match_speakers( + {"SPEAKER_00": (1.0, 0.0), "SPEAKER_01": (0.99, 0.01)}, prints, 0.5 + ) + + assert mapping == {"SPEAKER_00": "Coco", "SPEAKER_01": "Coco"} diff --git a/tests/unit/models/test_voice_print.py b/tests/unit/models/test_voice_print.py new file mode 100644 index 0000000..6ae0b21 --- /dev/null +++ b/tests/unit/models/test_voice_print.py @@ -0,0 +1,20 @@ +import pytest + +from vox.models.exceptions import ValidationError +from vox.models.voice_print import VoicePrint + + +class TestVoicePrint: + def test_create_when_valid_then_fields_set(self): + print_ = VoicePrint(name="Coco", embedding=(0.1, 0.2, 0.3)) + + assert print_.name == "Coco" + assert print_.embedding == (0.1, 0.2, 0.3) + + def test_create_when_blank_name_then_raises(self): + with pytest.raises(ValidationError, match="name"): + VoicePrint(name=" ", embedding=(0.1,)) + + def test_create_when_empty_embedding_then_raises(self): + with pytest.raises(ValidationError, match="embedding"): + VoicePrint(name="Coco", embedding=()) diff --git a/tests/unit/use_cases/test_enroll_speaker.py b/tests/unit/use_cases/test_enroll_speaker.py new file mode 100644 index 0000000..612ebbb --- /dev/null +++ b/tests/unit/use_cases/test_enroll_speaker.py @@ -0,0 +1,62 @@ +import pytest + +from tests.fakes.fake_voice_print_extractor import FakeVoicePrintExtractor +from tests.fakes.fake_voice_print_store import FakeVoicePrintStore +from vox.models.exceptions import ValidationError +from vox.use_cases.enroll_speaker import EnrollSpeakerRequest, EnrollSpeakerUseCase + + +def _request(**overrides) -> EnrollSpeakerRequest: + defaults = { + "name": "Coco", + "audio_path": "live.wav", + "start": 10.0, + "end": 25.0, + } + defaults.update(overrides) + return EnrollSpeakerRequest(**defaults) + + +class TestEnrollSpeaker: + def test_execute_when_valid_then_saves_print_under_name(self): + store = FakeVoicePrintStore() + use_case = EnrollSpeakerUseCase(FakeVoicePrintExtractor(), store) + + use_case.execute(_request()) + + assert len(store.saved) == 1 + assert store.saved[0].name == "Coco" + + def test_execute_when_valid_then_extracts_requested_span(self): + extractor = FakeVoicePrintExtractor() + use_case = EnrollSpeakerUseCase(extractor, FakeVoicePrintStore()) + + use_case.execute(_request(start=3.0, end=9.0)) + + assert extractor.extract_calls[0][1:] == (3.0, 9.0) + + def test_execute_when_end_before_start_then_raises(self): + use_case = EnrollSpeakerUseCase( + FakeVoicePrintExtractor(), FakeVoicePrintStore() + ) + + with pytest.raises(ValidationError, match="end"): + use_case.execute(_request(start=10.0, end=5.0)) + + def test_execute_when_span_too_short_then_raises(self): + use_case = EnrollSpeakerUseCase( + FakeVoicePrintExtractor(), FakeVoicePrintStore() + ) + + with pytest.raises(ValidationError, match="second"): + use_case.execute(_request(start=10.0, end=10.5)) + + def test_execute_when_valid_then_returns_name_and_span(self): + use_case = EnrollSpeakerUseCase( + FakeVoicePrintExtractor(), FakeVoicePrintStore() + ) + + response = use_case.execute(_request()) + + assert response.name == "Coco" + assert response.seconds_used == 15.0 diff --git a/tests/unit/use_cases/test_identify_speakers.py b/tests/unit/use_cases/test_identify_speakers.py new file mode 100644 index 0000000..1527891 --- /dev/null +++ b/tests/unit/use_cases/test_identify_speakers.py @@ -0,0 +1,60 @@ +from pathlib import Path + +from tests.fakes.fake_voice_print_extractor import FakeVoicePrintExtractor +from tests.fakes.fake_voice_print_store import FakeVoicePrintStore +from vox.models.speaker_turn import SpeakerTurn +from vox.models.voice_print import VoicePrint +from vox.use_cases.identify_speakers import IdentifySpeakersUseCase + +_TURNS = ( + SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=15.0, speaker="SPEAKER_01"), +) + + +class TestIdentifySpeakers: + def test_execute_when_no_prints_then_empty_mapping(self): + use_case = IdentifySpeakersUseCase( + extractor=FakeVoicePrintExtractor(), + store=FakeVoicePrintStore(), + ) + + assert use_case.execute(Path("a.wav"), _TURNS) == {} + + def test_execute_when_no_prints_then_extractor_not_called(self): + extractor = FakeVoicePrintExtractor() + use_case = IdentifySpeakersUseCase(extractor, FakeVoicePrintStore()) + + use_case.execute(Path("a.wav"), _TURNS) + + assert extractor.extract_calls == [] + + def test_execute_when_print_matches_then_label_mapped(self): + store = FakeVoicePrintStore((VoicePrint(name="Coco", embedding=(1.0, 0.0)),)) + use_case = IdentifySpeakersUseCase(FakeVoicePrintExtractor(), store) + + mapping = use_case.execute(Path("a.wav"), _TURNS) + + assert mapping == {"SPEAKER_00": "Coco", "SPEAKER_01": "Coco"} + + def test_execute_when_print_too_far_then_not_mapped(self): + store = FakeVoicePrintStore((VoicePrint(name="Coco", embedding=(0.0, 1.0)),)) + extractor = FakeVoicePrintExtractor({(0.0, 10.0): (1.0, 0.0)}) + use_case = IdentifySpeakersUseCase(extractor, store) + + mapping = use_case.execute(Path("a.wav"), _TURNS) + + assert "SPEAKER_00" not in mapping + + def test_execute_when_called_then_uses_longest_turn_per_speaker(self): + store = FakeVoicePrintStore((VoicePrint(name="Coco", embedding=(1.0, 0.0)),)) + extractor = FakeVoicePrintExtractor() + use_case = IdentifySpeakersUseCase(extractor, store) + turns = ( + SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), + SpeakerTurn(start=5.0, end=30.0, speaker="SPEAKER_00"), + ) + + use_case.execute(Path("a.wav"), turns) + + assert extractor.extract_calls == [(Path("a.wav"), 5.0, 30.0)] diff --git a/tests/unit/use_cases/test_transcribe.py b/tests/unit/use_cases/test_transcribe.py index b8dc732..2f99929 100644 --- a/tests/unit/use_cases/test_transcribe.py +++ b/tests/unit/use_cases/test_transcribe.py @@ -3,9 +3,11 @@ import pytest from tests.fakes.fake_audio_cleaner import FakeAudioCleaner +from tests.fakes.fake_diarizer import FakeDiarizer from tests.fakes.fake_downloader import FakeDownloader from tests.fakes.fake_file_writer import FakeFileWriter from tests.fakes.fake_progress import FakeProgressReporter +from tests.fakes.fake_speaker_identifier import FakeSpeakerIdentifier from tests.fakes.fake_transcriber import FakeTranscriber from vox.models.exceptions import ValidationError from vox.use_cases.transcribe import ( @@ -22,12 +24,16 @@ def __init__(self): self.transcriber = FakeTranscriber() self.file_writer = FakeFileWriter() self.progress = FakeProgressReporter() + self.diarizer = FakeDiarizer() + self.speaker_identifier = FakeSpeakerIdentifier() self.use_case = TranscribeUseCase( downloader=self.downloader, audio_cleaner=self.audio_cleaner, transcriber=self.transcriber, file_writer=self.file_writer, progress=self.progress, + diarizer=self.diarizer, + speaker_identifier=self.speaker_identifier, ) def execute(self, **overrides) -> TranscribeResponse: @@ -40,6 +46,8 @@ def execute(self, **overrides) -> TranscribeResponse: "no_clean": False, "no_download": False, "dry_run": False, + "diarize": False, + "identify": False, } defaults.update(overrides) request = TranscribeRequest(**defaults) @@ -198,3 +206,121 @@ def test_execute_when_url_input_then_filenames_use_vox_prefix(self): assert "vox_" in result.srt_path assert "vox_" in result.txt_path assert "vox_" in result.json_path + + +class TestExecuteWhenDiarizeThenAssignsSpeakers: + def test_execute_when_diarize_disabled_then_diarizer_not_called(self): + fix = TranscribeFixture() + + fix.execute(diarize=False) + + assert fix.diarizer.diarize_called_with is None + + def test_execute_when_diarize_then_calls_diarizer(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + assert fix.diarizer.diarize_called_with is not None + + def test_execute_when_diarize_then_diarizes_same_file_as_transcribed(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + transcribed_path = fix.transcriber.transcribe_called_with[0] + diarized_path = fix.diarizer.diarize_called_with[0] + assert diarized_path == transcribed_path + + def test_execute_when_diarize_then_written_segments_carry_speaker(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + written_result = fix.file_writer.json_written[0][0] + assert written_result.segments[0].speaker == "SPEAKER_00" + + def test_execute_when_diarize_disabled_then_segments_have_no_speaker(self): + fix = TranscribeFixture() + + fix.execute(diarize=False) + + written_result = fix.file_writer.json_written[0][0] + assert written_result.segments[0].speaker is None + + def test_execute_when_diarize_then_passes_num_speakers(self): + fix = TranscribeFixture() + + fix.execute(diarize=True, num_speakers=3) + + assert fix.diarizer.diarize_called_with[1] == 3 + + +class TestExecuteWhenDiarizeThenPreservesAudioTimeline: + def test_execute_when_diarize_then_silence_removal_disabled(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + config = fix.audio_cleaner.clean_called_with[1] + assert config.remove_silence is False + + def test_execute_when_diarize_then_denoise_disabled(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + config = fix.audio_cleaner.clean_called_with[1] + assert config.denoise is False + + def test_execute_when_diarize_then_still_mono_16k(self): + fix = TranscribeFixture() + + fix.execute(diarize=True) + + config = fix.audio_cleaner.clean_called_with[1] + assert (config.sample_rate, config.channels) == (16000, 1) + + def test_execute_when_no_diarize_then_default_cleaning_kept(self): + fix = TranscribeFixture() + + fix.execute(diarize=False) + + config = fix.audio_cleaner.clean_called_with[1] + assert config.remove_silence is True + assert config.denoise is True + + +class TestExecuteWhenIdentifyThenRenamesSpeakers: + def test_execute_when_identify_then_labels_replaced_by_names(self): + fix = TranscribeFixture() + fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} + + fix.execute(diarize=True, identify=True) + + written = fix.file_writer.json_written[0][0] + assert written.segments[0].speaker == "Coco" + + def test_execute_when_identify_disabled_then_labels_kept(self): + fix = TranscribeFixture() + fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} + + fix.execute(diarize=True, identify=False) + + written = fix.file_writer.json_written[0][0] + assert written.segments[0].speaker == "SPEAKER_00" + + def test_execute_when_identify_without_diarize_then_no_identification(self): + fix = TranscribeFixture() + + fix.execute(diarize=False, identify=True) + + assert fix.speaker_identifier.called_with is None + + def test_execute_when_identify_then_receives_diarized_turns(self): + fix = TranscribeFixture() + + fix.execute(diarize=True, identify=True) + + _path, turns = fix.speaker_identifier.called_with + assert turns[0].speaker == "SPEAKER_00" diff --git a/uv.lock b/uv.lock index abaf654..a58b807 100644 --- a/uv.lock +++ b/uv.lock @@ -979,6 +979,45 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e0/f9/0595336914c5619e5f28a1fb793285925a8cd4b432c9da0a987836c7f822/shellingham-1.5.4-py2.py3-none-any.whl", hash = "sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686", size = 9755, upload-time = "2023-10-24T04:13:38.866Z" }, ] +[[package]] +name = "sherpa-onnx" +version = "1.13.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/ca/f8/735244770b4bc63f85fabdad0e46d6ec1f4cc24e64f6e082c2e0fea92b8c/sherpa_onnx-1.13.4.tar.gz", hash = "sha256:29547692418513ad88034c2b5f98985e33042b2351e4ab375469f19a8de18c5f", size = 982750, upload-time = "2026-07-07T13:04:55.145Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4b/36/45b17335f041f1383f6fd142ab57c2d8a337ba2386b7547b125ec9d780af/sherpa_onnx-1.13.4-cp313-cp313-linux_armv7l.whl", hash = "sha256:9e98dc5e0559ad953f227fc884958c71b10c65a93667331405e7d4441ed5f76d", size = 11932654, upload-time = "2026-07-07T14:18:59.791Z" }, + { url = "https://files.pythonhosted.org/packages/68/d7/1e9a7dedab2da8af1a8417b4f4d5f496bd7700a71b59c5e085de5e10761b/sherpa_onnx-1.13.4-cp313-cp313-macosx_10_15_universal2.whl", hash = "sha256:083747c2d0362ead0501cc773a618be19862a800b3f8f259d3bd3486f1494af4", size = 4372494, upload-time = "2026-07-07T13:02:05.512Z" }, + { url = "https://files.pythonhosted.org/packages/62/28/09aa9461e8bdf894ba8466e047e40fb5da8aaa6d68c19cf1e2aabe01e706/sherpa_onnx-1.13.4-cp313-cp313-macosx_10_15_x86_64.whl", hash = "sha256:b4f54363b264b16148a724b4442f00cda97fcd4e9beeda3d75637753910e8557", size = 2307539, upload-time = "2026-07-07T12:28:22.056Z" }, + { url = "https://files.pythonhosted.org/packages/21/ed/d07787dd4be4119e6587c840f6b417c2d57c14d694d334af609d68cb5a41/sherpa_onnx-1.13.4-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:8ec5394b4ea73bf01e6883cf078348f87350f4eb3567d51d92cae77ea2582403", size = 2115586, upload-time = "2026-07-07T12:47:31.367Z" }, + { url = "https://files.pythonhosted.org/packages/6d/c2/281d84dc9e448ea99d7fb77708cbe1cc7cfd8c7d669727dc94385a9e4ca5/sherpa_onnx-1.13.4-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:a39352ceb2ec6671a1f252fb768fff75bb2f0bc849cca5f66f490e89910a860d", size = 4136268, upload-time = "2026-07-07T12:10:09.762Z" }, + { url = "https://files.pythonhosted.org/packages/db/47/da3ea14ab647a4f6580227853fe29353e1173ff77064d42c0bb31d01b453/sherpa_onnx-1.13.4-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:88af596be24eac32982dd64fcac30af99d9130ca498bfa1a0064189c8498195b", size = 4358385, upload-time = "2026-07-07T13:10:10.265Z" }, + { url = "https://files.pythonhosted.org/packages/44/90/84205ff383ba9335c3821c2cd6d514350f52199cb640ec094517f1f911a0/sherpa_onnx-1.13.4-cp313-cp313-win32.whl", hash = "sha256:0cabb508a15be22138f9fb7695d7ec5f3893ecd088ee419b9df559fed7e8f649", size = 1929707, upload-time = "2026-07-07T12:51:53.196Z" }, + { url = "https://files.pythonhosted.org/packages/82/40/ee8a0a8c83fc6d7f5245a5a031e471d3b115e20cce867e7abb2f9d4185c9/sherpa_onnx-1.13.4-cp313-cp313-win_amd64.whl", hash = "sha256:17050fdfb48d37ae996364f697c554a1399740d18e5a56b143c011d00cfed3e0", size = 2244504, upload-time = "2026-07-07T12:32:59.882Z" }, + { url = "https://files.pythonhosted.org/packages/8e/9e/cb97cf04c0c4de0a1d43463952e387a938eea8f6ab94e2eb95af21162f57/sherpa_onnx-1.13.4-cp314-cp314-linux_armv7l.whl", hash = "sha256:7ddb3d46fe0d6cc745d222e055844dbb24fd4e66935d4ead67d1315911ea7a46", size = 11930841, upload-time = "2026-07-07T14:37:28.79Z" }, + { url = "https://files.pythonhosted.org/packages/cc/0f/7fbed45ef8437d20967f4577514e8903b7f4a87a8cea9b9fed9e2b18fa45/sherpa_onnx-1.13.4-cp314-cp314-macosx_10_15_universal2.whl", hash = "sha256:ff9237b98c173dabf8f5f6317c4ead54afd4b35179b03be652f65e1e121bee74", size = 4425178, upload-time = "2026-07-07T13:01:45.126Z" }, + { url = "https://files.pythonhosted.org/packages/9a/90/191b3b73af1b54584f9f96e256d9ca9f6f3781cacdb4287c7efea6b25094/sherpa_onnx-1.13.4-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:972137cbf19d501ea6a51857528b340ce1c3d204571e3d56dcc72eded5827418", size = 2345420, upload-time = "2026-07-07T11:46:13.903Z" }, + { url = "https://files.pythonhosted.org/packages/1f/b5/29997d7de29cae8e3f54e8a993fcf14d48d5a7960deab10df479fcbc7d64/sherpa_onnx-1.13.4-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:ff38ac153527fda0f6158e7c097b66e980a6840455d908e037ce09a4a4c1ce14", size = 2118562, upload-time = "2026-07-07T12:26:40.457Z" }, + { url = "https://files.pythonhosted.org/packages/8c/e5/8c601626448358ac8571100db051f35af5ef02c1d12ab1cb026436474a22/sherpa_onnx-1.13.4-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:7a6bf37e45727b0e94661a29b34927e9edc4b3849411cc38036e30d81df48d31", size = 4142540, upload-time = "2026-07-07T12:03:40.699Z" }, + { url = "https://files.pythonhosted.org/packages/5a/4b/49b4be95af2e12bfa8a92a7eba720cc68e4da178847ffa4ddb66479d4e9c/sherpa_onnx-1.13.4-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:16b046f1c7ceecf666947d652c928278c377847db7f697c9e94af9feaf20ac2c", size = 4360312, upload-time = "2026-07-07T13:05:36.871Z" }, + { url = "https://files.pythonhosted.org/packages/96/1c/a80bfad89846ccb1da4d18ef46bf364c83cff9514a0e298110c8de480e14/sherpa_onnx-1.13.4-cp314-cp314-win32.whl", hash = "sha256:39af016b9fb7c053da2270e44182fb8932eeab991dcdc963bbc4366324e9ec51", size = 1970019, upload-time = "2026-07-07T12:13:06.663Z" }, + { url = "https://files.pythonhosted.org/packages/89/02/1e10f71be635a3f9ef793b07ab88ddb7afa82ab722f3ea275a4877da1bcb/sherpa_onnx-1.13.4-cp314-cp314-win_amd64.whl", hash = "sha256:cb1834182c4047b8edb1dceeed8d5cf7d6e10295a4079e5e0fea674b4314db06", size = 2308592, upload-time = "2026-07-07T13:03:09.057Z" }, +] + +[[package]] +name = "sherpa-onnx-core" +version = "1.13.4" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ad/69/ae33a8cb1ecc30e0c4638d76772e56161162d57bef748380790e1257841f/sherpa_onnx_core-1.13.4-py3-none-macosx_10_15_universal2.whl", hash = "sha256:737099a817998e4d74379dcd44d8d1b332a3fa1822780be174334c3d1d1e2451", size = 36025264, upload-time = "2026-07-07T11:41:37.506Z" }, + { url = "https://files.pythonhosted.org/packages/2a/8c/a1336beab226d228f62bdbe1cacf1439a2f6cf8714baac98f1f031dd2f60/sherpa_onnx_core-1.13.4-py3-none-macosx_10_15_x86_64.whl", hash = "sha256:674ea57eb6458002dab4e39749d21a7364b2f962f4a4958a3ee4c351b093cafd", size = 19325474, upload-time = "2026-07-07T11:31:03.394Z" }, + { url = "https://files.pythonhosted.org/packages/a2/26/d0d6bebea4ef8b7de6eed1a335235d1acce9197425c08eecb57ef107904d/sherpa_onnx_core-1.13.4-py3-none-macosx_11_0_arm64.whl", hash = "sha256:7820183581b711a68e30281a4fb36c0af8ef5615bfc789234fb259439824a014", size = 16897650, upload-time = "2026-07-07T11:42:28.755Z" }, + { url = "https://files.pythonhosted.org/packages/81/29/15859e896574230d5377738ef27e8647cbd19fa325d75b1a00074791798d/sherpa_onnx_core-1.13.4-py3-none-manylinux2014_aarch64.whl", hash = "sha256:b4e4d17eb0d5c569bf4c9effcbc3daef57cf7e2b8418e07ae90f90c9b60b35d5", size = 13023017, upload-time = "2026-07-07T11:41:30.935Z" }, + { url = "https://files.pythonhosted.org/packages/41/be/38c57721d71ee74d984b1ca21720a8ca8477d6d341026af24ff658866ef9/sherpa_onnx_core-1.13.4-py3-none-manylinux2014_x86_64.whl", hash = "sha256:367aa06cee90b3fd7959d4e071d6fc821710b859af399b4987e5c3119ee6ae2a", size = 10338839, upload-time = "2026-07-07T12:21:33.627Z" }, + { url = "https://files.pythonhosted.org/packages/7e/17/949ee6d5cff4ec9e2dcda014276a798b814e2eb80d35d820884e6e0614ff/sherpa_onnx_core-1.13.4-py3-none-manylinux_2_35_armv7l.whl", hash = "sha256:3b9cc7da5ce4a2333a9ae211f39c34dee981b13f2c9c8680fbf2d9636e87d617", size = 9889942, upload-time = "2026-07-07T14:25:04.94Z" }, + { url = "https://files.pythonhosted.org/packages/3b/71/e2fb4b965b86bd3ee6ca35d81999e9edad47c20d2c8fc4e03db8083494bc/sherpa_onnx_core-1.13.4-py3-none-win32.whl", hash = "sha256:ba9de2f463ae67ee2947f1bd9b981a1e39e6b831e281bf5088068251b8ec4dde", size = 13832564, upload-time = "2026-07-07T11:49:11.781Z" }, + { url = "https://files.pythonhosted.org/packages/95/b0/c3d59ac76f3db873e41bd0cb4fc30b352a278da3289217985aaae3650211/sherpa_onnx_core-1.13.4-py3-none-win_amd64.whl", hash = "sha256:0a6949cf0fd83adb9fbcfdf5c27b8907a57f7b48626db703c7f6037be9b61764", size = 16450053, upload-time = "2026-07-07T12:03:05.196Z" }, +] + [[package]] name = "sympy" version = "1.14.0" @@ -1161,6 +1200,8 @@ source = { editable = "." } dependencies = [ { name = "click" }, { name = "mlx-whisper" }, + { name = "sherpa-onnx" }, + { name = "sherpa-onnx-core" }, { name = "yt-dlp" }, ] @@ -1175,6 +1216,8 @@ dev = [ requires-dist = [ { name = "click", specifier = ">=8.1" }, { name = "mlx-whisper", specifier = ">=0.4" }, + { name = "sherpa-onnx", specifier = ">=1.13.4" }, + { name = "sherpa-onnx-core", specifier = ">=1.13.4" }, { name = "yt-dlp", specifier = ">=2024.0.0" }, ] From be5a70204e90a4b84f1ff971684a0317d3c2d9fc Mon Sep 17 00:00:00 2001 From: Yoann Date: Fri, 31 Jul 2026 16:06:48 +0200 Subject: [PATCH 2/5] fix: identify known speakers without asking for a flag --identify was pointless: when no voice is enrolled the use case returns early without extracting a single embedding, so running identification unconditionally costs nothing. Whenever voices are known, matching speakers are now renamed automatically. Replaced by --no-identify for the rare case of wanting raw SPEAKER_xx labels back. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Uz3rnGtYHs6YcBNr6MkY1z --- src/vox/adapters/cli/transcribe_cmd.py | 8 ++++---- src/vox/schemas/transcribe.json | 4 ++-- src/vox/use_cases/transcribe.py | 4 ++-- tests/unit/use_cases/test_transcribe.py | 16 ++++++++-------- 4 files changed, 16 insertions(+), 16 deletions(-) diff --git a/src/vox/adapters/cli/transcribe_cmd.py b/src/vox/adapters/cli/transcribe_cmd.py index 32da7db..769c54a 100644 --- a/src/vox/adapters/cli/transcribe_cmd.py +++ b/src/vox/adapters/cli/transcribe_cmd.py @@ -50,9 +50,9 @@ ) @click.option("--diarize", is_flag=True, help="Identify speakers (who said what)") @click.option( - "--identify", + "--no-identify", is_flag=True, - help="Replace SPEAKER_xx by known names (see: vox speakers add)", + help="Keep SPEAKER_xx labels even when voices are known", ) @click.option( "--speakers", @@ -76,7 +76,7 @@ def transcribe( browser, backend, diarize, - identify, + no_identify, speakers, ): source, language, model = _apply_json_overrides( @@ -100,7 +100,7 @@ def transcribe( no_download=no_download, dry_run=dry_run, diarize=diarize, - identify=identify, + no_identify=no_identify, num_speakers=speakers, ) try: diff --git a/src/vox/schemas/transcribe.json b/src/vox/schemas/transcribe.json index 1a64891..f7d6757 100644 --- a/src/vox/schemas/transcribe.json +++ b/src/vox/schemas/transcribe.json @@ -67,10 +67,10 @@ "type": "integer", "description": "Known speaker count. Auto-detected when omitted, but passing the exact count is markedly more reliable." }, - "identify": { + "no-identify": { "type": "boolean", "default": false, - "description": "Replace SPEAKER_00 labels with real names, matched against voices registered via 'vox speakers add'. Requires diarize." + "description": "Keep raw SPEAKER_00 labels. By default, whenever voices have been registered via 'vox speakers add', matching speakers are renamed automatically." } }, "output": { diff --git a/src/vox/use_cases/transcribe.py b/src/vox/use_cases/transcribe.py index 49a320a..4aab980 100644 --- a/src/vox/use_cases/transcribe.py +++ b/src/vox/use_cases/transcribe.py @@ -31,7 +31,7 @@ class TranscribeRequest: dry_run: bool = False output_stem: str = "" diarize: bool = False - identify: bool = False + no_identify: bool = False num_speakers: int | None = None @@ -120,7 +120,7 @@ def _maybe_diarize(self, result, audio_path: Path, request): return replace(result, segments=segments) def _maybe_identify(self, segments, turns, audio_path: Path, request): - if not request.identify or self._speaker_identifier is None: + if request.no_identify or self._speaker_identifier is None: return segments self._progress.update("Matching known voices") mapping = self._speaker_identifier.execute(audio_path, turns) diff --git a/tests/unit/use_cases/test_transcribe.py b/tests/unit/use_cases/test_transcribe.py index 2f99929..12a4883 100644 --- a/tests/unit/use_cases/test_transcribe.py +++ b/tests/unit/use_cases/test_transcribe.py @@ -47,7 +47,7 @@ def execute(self, **overrides) -> TranscribeResponse: "no_download": False, "dry_run": False, "diarize": False, - "identify": False, + "no_identify": False, } defaults.update(overrides) request = TranscribeRequest(**defaults) @@ -292,35 +292,35 @@ def test_execute_when_no_diarize_then_default_cleaning_kept(self): class TestExecuteWhenIdentifyThenRenamesSpeakers: - def test_execute_when_identify_then_labels_replaced_by_names(self): + def test_execute_when_voices_known_then_labels_replaced_without_any_flag(self): fix = TranscribeFixture() fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} - fix.execute(diarize=True, identify=True) + fix.execute(diarize=True) written = fix.file_writer.json_written[0][0] assert written.segments[0].speaker == "Coco" - def test_execute_when_identify_disabled_then_labels_kept(self): + def test_execute_when_no_identify_then_labels_kept(self): fix = TranscribeFixture() fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} - fix.execute(diarize=True, identify=False) + fix.execute(diarize=True, no_identify=True) written = fix.file_writer.json_written[0][0] assert written.segments[0].speaker == "SPEAKER_00" - def test_execute_when_identify_without_diarize_then_no_identification(self): + def test_execute_when_no_diarize_then_no_identification(self): fix = TranscribeFixture() - fix.execute(diarize=False, identify=True) + fix.execute(diarize=False) assert fix.speaker_identifier.called_with is None def test_execute_when_identify_then_receives_diarized_turns(self): fix = TranscribeFixture() - fix.execute(diarize=True, identify=True) + fix.execute(diarize=True) _path, turns = fix.speaker_identifier.called_with assert turns[0].speaker == "SPEAKER_00" From fdfd889301e966b7f5af8e7616f0eb5ae5301011 Mon Sep 17 00:00:00 2001 From: Yoann Date: Fri, 31 Jul 2026 16:18:38 +0200 Subject: [PATCH 3/5] feat: detect the speaker count automatically --speakers is no longer needed. The diarizer now over-segments on purpose (clustering threshold 0.05), embeds the longest turn of each resulting label, and merges the excess labels by picking the count with the best silhouette score. Measured on two-voice dialogues: both are resolved to exactly 2 speakers with the right alternation, including the pair of near-identical synthesised voices that a single fixed threshold could never separate. - clustering, silhouette, estimate_speaker_count, merge_labels and relabel_turns are pure domain, plain Python, no numpy and no sherpa - AutoSpeakerCountDiarizer implements the Diarizer port by composition, so it is unit-tested entirely with fakes - cost stays low: one embedding per label (a handful), not per turn (628 on a one-hour file) - --speakers still forces a count and then skips estimation entirely Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Uz3rnGtYHs6YcBNr6MkY1z --- .../adapters/auto_speaker_count_diarizer.py | 51 ++++++++++ src/vox/adapters/cli/transcribe_cmd.py | 11 ++- src/vox/models/clustering.py | 45 +++++++++ src/vox/models/silhouette.py | 58 +++++++++++ src/vox/models/speaker_count.py | 39 ++++++++ src/vox/models/turn_relabeling.py | 28 ++++++ src/vox/schemas/transcribe.json | 2 +- .../test_auto_speaker_count_diarizer.py | 97 +++++++++++++++++++ tests/unit/models/test_clustering.py | 52 ++++++++++ tests/unit/models/test_silhouette.py | 42 ++++++++ tests/unit/models/test_speaker_count.py | 41 ++++++++ tests/unit/models/test_turn_relabeling.py | 45 +++++++++ 12 files changed, 508 insertions(+), 3 deletions(-) create mode 100644 src/vox/adapters/auto_speaker_count_diarizer.py create mode 100644 src/vox/models/clustering.py create mode 100644 src/vox/models/silhouette.py create mode 100644 src/vox/models/speaker_count.py create mode 100644 src/vox/models/turn_relabeling.py create mode 100644 tests/unit/adapters/test_auto_speaker_count_diarizer.py create mode 100644 tests/unit/models/test_clustering.py create mode 100644 tests/unit/models/test_silhouette.py create mode 100644 tests/unit/models/test_speaker_count.py create mode 100644 tests/unit/models/test_turn_relabeling.py diff --git a/src/vox/adapters/auto_speaker_count_diarizer.py b/src/vox/adapters/auto_speaker_count_diarizer.py new file mode 100644 index 0000000..a8403d8 --- /dev/null +++ b/src/vox/adapters/auto_speaker_count_diarizer.py @@ -0,0 +1,51 @@ +from pathlib import Path + +from vox.models.clustering import cluster_embeddings +from vox.models.speaker_count import DEFAULT_MAX_SPEAKERS, estimate_speaker_count +from vox.models.speaker_turn import SpeakerTurn +from vox.models.turn_relabeling import merge_labels, relabel_turns +from vox.models.turn_selection import longest_turn_per_speaker +from vox.ports.diarizer import Diarizer +from vox.ports.voice_print_extractor import VoicePrintExtractor + +OVERSEGMENTATION_THRESHOLD = 0.05 + + +class AutoSpeakerCountDiarizer: + def __init__( + self, + diarizer: Diarizer, + extractor: VoicePrintExtractor, + max_speakers: int = DEFAULT_MAX_SPEAKERS, + ): + self._diarizer = diarizer + self._extractor = extractor + self._max_speakers = max_speakers + + def diarize( + self, + audio_path: Path, + num_speakers: int | None, + ) -> tuple[SpeakerTurn, ...]: + if num_speakers: + return self._diarizer.diarize(audio_path, num_speakers) + turns = self._diarizer.diarize(audio_path, None) + longest = longest_turn_per_speaker(turns) + if len(longest) < 2: + return turns + return self._merge_over_segmented(turns, longest, audio_path) + + def _merge_over_segmented( + self, + turns: tuple[SpeakerTurn, ...], + longest: dict[str, SpeakerTurn], + audio_path: Path, + ) -> tuple[SpeakerTurn, ...]: + labels = list(longest) + embeddings = [ + self._extractor.extract(audio_path, longest[a].start, longest[a].end) + for a in labels + ] + count = estimate_speaker_count(embeddings, self._max_speakers) + groups = cluster_embeddings(embeddings, count) + return relabel_turns(turns, merge_labels(labels, groups)) diff --git a/src/vox/adapters/cli/transcribe_cmd.py b/src/vox/adapters/cli/transcribe_cmd.py index 769c54a..fa2e84a 100644 --- a/src/vox/adapters/cli/transcribe_cmd.py +++ b/src/vox/adapters/cli/transcribe_cmd.py @@ -3,6 +3,10 @@ import click +from vox.adapters.auto_speaker_count_diarizer import ( + OVERSEGMENTATION_THRESHOLD, + AutoSpeakerCountDiarizer, +) from vox.adapters.cli.open_hint import format_open_hint from vox.adapters.cli.output_formatter import format_output from vox.adapters.click_progress import ClickProgressReporter @@ -58,7 +62,7 @@ "--speakers", type=int, default=None, - help="Known speaker count (auto-detected if omitted)", + help="Force the speaker count (detected automatically if omitted)", ) def transcribe( source, @@ -143,7 +147,10 @@ def _build_use_case( transcriber=_build_transcriber(backend), file_writer=DiskFileWriter(), progress=ClickProgressReporter(), - diarizer=SherpaDiarizer(), + diarizer=AutoSpeakerCountDiarizer( + SherpaDiarizer(clustering_threshold=OVERSEGMENTATION_THRESHOLD), + SherpaVoicePrintExtractor(), + ), speaker_identifier=IdentifySpeakersUseCase( extractor=SherpaVoicePrintExtractor(), store=JsonVoicePrintStore(), diff --git a/src/vox/models/clustering.py b/src/vox/models/clustering.py new file mode 100644 index 0000000..2ebb744 --- /dev/null +++ b/src/vox/models/clustering.py @@ -0,0 +1,45 @@ +from vox.models.voice_matching import Embedding, cosine_similarity + + +def cluster_embeddings( + embeddings: list[Embedding], + n_clusters: int, +) -> list[int]: + groups = [[i] for i in range(len(embeddings))] + target = max(1, min(n_clusters, len(groups))) + while len(groups) > target: + left, right = _closest_pair(groups, embeddings) + groups[left].extend(groups.pop(right)) + return _to_labels(groups, len(embeddings)) + + +def _closest_pair( + groups: list[list[int]], + embeddings: list[Embedding], +) -> tuple[int, int]: + best = (-2.0, 0, 1) + for i in range(len(groups)): + for j in range(i + 1, len(groups)): + score = _average_linkage(groups[i], groups[j], embeddings) + if score > best[0]: + best = (score, i, j) + return best[1], best[2] + + +def _average_linkage( + left: list[int], + right: list[int], + embeddings: list[Embedding], +) -> float: + scores = [ + cosine_similarity(embeddings[a], embeddings[b]) for a in left for b in right + ] + return sum(scores) / len(scores) + + +def _to_labels(groups: list[list[int]], size: int) -> list[int]: + labels = [0] * size + for label, group in enumerate(groups): + for index in group: + labels[index] = label + return labels diff --git a/src/vox/models/silhouette.py b/src/vox/models/silhouette.py new file mode 100644 index 0000000..e426d0f --- /dev/null +++ b/src/vox/models/silhouette.py @@ -0,0 +1,58 @@ +from vox.models.voice_matching import Embedding, cosine_similarity + + +def silhouette_score( + embeddings: list[Embedding], + labels: list[int], +) -> float: + clusters = _group_indices(labels) + if len(clusters) < 2: + return 0.0 + scores = [ + _point_score(index, embeddings, clusters, labels) + for index in range(len(embeddings)) + ] + return sum(scores) / len(scores) + + +def _point_score( + index: int, + embeddings: list[Embedding], + clusters: dict[int, list[int]], + labels: list[int], +) -> float: + own = clusters[labels[index]] + if len(own) < 2: + return 0.0 + inside = _mean_distance(index, own, embeddings, exclude_self=True) + outside = min( + _mean_distance(index, members, embeddings, exclude_self=False) + for label, members in clusters.items() + if label != labels[index] + ) + spread = max(inside, outside) + if spread == 0: + return 0.0 + return (outside - inside) / spread + + +def _mean_distance( + index: int, + members: list[int], + embeddings: list[Embedding], + exclude_self: bool, +) -> float: + others = [m for m in members if m != index] if exclude_self else members + distances = [_distance(embeddings[index], embeddings[m]) for m in others] + return sum(distances) / len(distances) + + +def _distance(left: Embedding, right: Embedding) -> float: + return 1.0 - cosine_similarity(left, right) + + +def _group_indices(labels: list[int]) -> dict[int, list[int]]: + clusters: dict[int, list[int]] = {} + for index, label in enumerate(labels): + clusters.setdefault(label, []).append(index) + return clusters diff --git a/src/vox/models/speaker_count.py b/src/vox/models/speaker_count.py new file mode 100644 index 0000000..07a0333 --- /dev/null +++ b/src/vox/models/speaker_count.py @@ -0,0 +1,39 @@ +from vox.models.clustering import cluster_embeddings +from vox.models.silhouette import silhouette_score +from vox.models.voice_matching import Embedding, cosine_similarity + +DEFAULT_MAX_SPEAKERS = 10 +SAME_VOICE_SIMILARITY = 0.85 + + +def estimate_speaker_count( + embeddings: list[Embedding], + max_speakers: int = DEFAULT_MAX_SPEAKERS, +) -> int: + if len(embeddings) < 2: + return 1 + if _all_one_voice(embeddings): + return 1 + return _best_scoring_count(embeddings, max_speakers) + + +def _best_scoring_count( + embeddings: list[Embedding], + max_speakers: int, +) -> int: + ceiling = min(max_speakers, len(embeddings)) + scored = [ + (silhouette_score(embeddings, cluster_embeddings(embeddings, n)), n) + for n in range(2, ceiling + 1) + ] + if not scored: + return 1 + return max(scored)[1] + + +def _all_one_voice(embeddings: list[Embedding]) -> bool: + return all( + cosine_similarity(embeddings[i], embeddings[j]) >= SAME_VOICE_SIMILARITY + for i in range(len(embeddings)) + for j in range(i + 1, len(embeddings)) + ) diff --git a/src/vox/models/turn_relabeling.py b/src/vox/models/turn_relabeling.py new file mode 100644 index 0000000..c1eb586 --- /dev/null +++ b/src/vox/models/turn_relabeling.py @@ -0,0 +1,28 @@ +from dataclasses import replace + +from vox.models.speaker_turn import SpeakerTurn + + +def relabel_turns( + turns: tuple[SpeakerTurn, ...], + mapping: dict[str, str], +) -> tuple[SpeakerTurn, ...]: + if not mapping: + return turns + return tuple(replace(t, speaker=mapping.get(t.speaker, t.speaker)) for t in turns) + + +def merge_labels(labels: list[str], groups: list[int]) -> dict[str, str]: + renumbered = _renumber(groups) + return { + label: f"SPEAKER_{renumbered[group]:02d}" + for label, group in zip(labels, groups, strict=True) + } + + +def _renumber(groups: list[int]) -> dict[int, int]: + order: dict[int, int] = {} + for group in groups: + if group not in order: + order[group] = len(order) + return order diff --git a/src/vox/schemas/transcribe.json b/src/vox/schemas/transcribe.json index f7d6757..2e23626 100644 --- a/src/vox/schemas/transcribe.json +++ b/src/vox/schemas/transcribe.json @@ -65,7 +65,7 @@ }, "speakers": { "type": "integer", - "description": "Known speaker count. Auto-detected when omitted, but passing the exact count is markedly more reliable." + "description": "Force the speaker count. Detected automatically when omitted, by over-segmenting then merging labels on a silhouette score; pass it only to override a wrong estimate." }, "no-identify": { "type": "boolean", diff --git a/tests/unit/adapters/test_auto_speaker_count_diarizer.py b/tests/unit/adapters/test_auto_speaker_count_diarizer.py new file mode 100644 index 0000000..f02ee86 --- /dev/null +++ b/tests/unit/adapters/test_auto_speaker_count_diarizer.py @@ -0,0 +1,97 @@ +from pathlib import Path + +from tests.fakes.fake_diarizer import FakeDiarizer +from tests.fakes.fake_voice_print_extractor import FakeVoicePrintExtractor +from vox.adapters.auto_speaker_count_diarizer import AutoSpeakerCountDiarizer +from vox.models.speaker_turn import SpeakerTurn + +_AUDIO = Path("live.wav") + +_THREE_LABELS = ( + SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=20.0, speaker="SPEAKER_01"), + SpeakerTurn(start=20.0, end=30.0, speaker="SPEAKER_02"), +) + + +def _labels(turns) -> list[str]: + return [t.speaker for t in turns] + + +class TestAutoSpeakerCountDiarizer: + def test_diarize_when_count_forced_then_delegates_untouched(self): + inner = FakeDiarizer(_THREE_LABELS) + auto = AutoSpeakerCountDiarizer(inner, FakeVoicePrintExtractor()) + + auto.diarize(_AUDIO, 3) + + assert inner.diarize_called_with == (_AUDIO, 3) + + def test_diarize_when_count_forced_then_no_embedding_computed(self): + extractor = FakeVoicePrintExtractor() + auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + + auto.diarize(_AUDIO, 3) + + assert extractor.extract_calls == [] + + def test_diarize_when_single_label_then_returned_as_is(self): + turns = (SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"),) + auto = AutoSpeakerCountDiarizer(FakeDiarizer(turns), FakeVoicePrintExtractor()) + + assert auto.diarize(_AUDIO, None) == turns + + def test_diarize_when_labels_are_distinct_voices_then_all_kept(self): + extractor = FakeVoicePrintExtractor( + { + (0.0, 10.0): (1.0, 0.0, 0.0), + (10.0, 20.0): (0.0, 1.0, 0.0), + (20.0, 30.0): (0.0, 0.0, 1.0), + } + ) + auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + + assert len(set(_labels(auto.diarize(_AUDIO, None)))) == 3 + + def test_diarize_when_two_labels_are_the_same_voice_then_merged(self): + extractor = FakeVoicePrintExtractor( + { + (0.0, 10.0): (1.0, 0.0), + (10.0, 20.0): (0.0, 1.0), + (20.0, 30.0): (0.99, 0.01), + } + ) + auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + + result = _labels(auto.diarize(_AUDIO, None)) + + assert len(set(result)) == 2 + assert result[0] == result[2] + assert result[0] != result[1] + + def test_diarize_when_all_labels_same_voice_then_collapses_to_one(self): + extractor = FakeVoicePrintExtractor( + { + (0.0, 10.0): (1.0, 0.0), + (10.0, 20.0): (0.999, 0.01), + (20.0, 30.0): (0.998, 0.02), + } + ) + auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + + assert len(set(_labels(auto.diarize(_AUDIO, None)))) == 1 + + def test_diarize_when_auto_then_uses_longest_turn_of_each_label(self): + extractor = FakeVoicePrintExtractor() + turns = ( + SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), + SpeakerTurn(start=5.0, end=30.0, speaker="SPEAKER_00"), + SpeakerTurn(start=30.0, end=40.0, speaker="SPEAKER_01"), + ) + auto = AutoSpeakerCountDiarizer(FakeDiarizer(turns), extractor) + + auto.diarize(_AUDIO, None) + + spans = [call[1:] for call in extractor.extract_calls] + assert (5.0, 30.0) in spans + assert (0.0, 2.0) not in spans diff --git a/tests/unit/models/test_clustering.py b/tests/unit/models/test_clustering.py new file mode 100644 index 0000000..3d7240b --- /dev/null +++ b/tests/unit/models/test_clustering.py @@ -0,0 +1,52 @@ +from vox.models.clustering import cluster_embeddings + +_TWO_GROUPS = [ + (1.0, 0.0), + (0.99, 0.01), + (0.0, 1.0), + (0.01, 0.99), +] + + +def _groups(labels: list[int]) -> set[frozenset[int]]: + groups: dict[int, set[int]] = {} + for index, label in enumerate(labels): + groups.setdefault(label, set()).add(index) + return {frozenset(g) for g in groups.values()} + + +class TestClusterEmbeddings: + def test_cluster_when_two_obvious_groups_then_splits_them(self): + labels = cluster_embeddings(_TWO_GROUPS, n_clusters=2) + + assert _groups(labels) == {frozenset({0, 1}), frozenset({2, 3})} + + def test_cluster_when_one_cluster_then_all_together(self): + labels = cluster_embeddings(_TWO_GROUPS, n_clusters=1) + + assert len(set(labels)) == 1 + + def test_cluster_when_as_many_clusters_as_points_then_all_separate(self): + labels = cluster_embeddings(_TWO_GROUPS, n_clusters=4) + + assert len(set(labels)) == 4 + + def test_cluster_when_more_clusters_than_points_then_capped(self): + labels = cluster_embeddings(_TWO_GROUPS, n_clusters=99) + + assert len(set(labels)) == 4 + + def test_cluster_when_empty_then_empty(self): + assert cluster_embeddings([], n_clusters=2) == [] + + def test_cluster_when_three_groups_then_splits_them(self): + points = [(1.0, 0.0, 0.0), (0.0, 1.0, 0.0), (0.0, 0.0, 1.0)] + + labels = cluster_embeddings(points, n_clusters=3) + + assert len(set(labels)) == 3 + + def test_cluster_when_called_then_labels_are_contiguous_from_zero(self): + labels = cluster_embeddings(_TWO_GROUPS, n_clusters=2) + + assert sorted(set(labels)) == [0, 1] diff --git a/tests/unit/models/test_silhouette.py b/tests/unit/models/test_silhouette.py new file mode 100644 index 0000000..bb83c99 --- /dev/null +++ b/tests/unit/models/test_silhouette.py @@ -0,0 +1,42 @@ +from vox.models.silhouette import silhouette_score + +_WELL_SEPARATED = [ + (1.0, 0.0), + (0.99, 0.01), + (0.0, 1.0), + (0.01, 0.99), +] + + +class TestSilhouetteScore: + def test_score_when_well_separated_then_close_to_one(self): + score = silhouette_score(_WELL_SEPARATED, [0, 0, 1, 1]) + + assert score > 0.9 + + def test_score_when_groups_are_scrambled_then_much_lower(self): + good = silhouette_score(_WELL_SEPARATED, [0, 0, 1, 1]) + scrambled = silhouette_score(_WELL_SEPARATED, [0, 1, 0, 1]) + + assert scrambled < good + + def test_score_when_single_cluster_then_zero(self): + assert silhouette_score(_WELL_SEPARATED, [0, 0, 0, 0]) == 0.0 + + def test_score_when_every_point_alone_then_zero(self): + assert silhouette_score(_WELL_SEPARATED, [0, 1, 2, 3]) == 0.0 + + def test_score_when_empty_then_zero(self): + assert silhouette_score([], []) == 0.0 + + def test_score_when_three_clear_groups_then_high(self): + points = [ + (1.0, 0.0, 0.0), + (0.99, 0.01, 0.0), + (0.0, 1.0, 0.0), + (0.0, 0.99, 0.01), + (0.0, 0.0, 1.0), + (0.01, 0.0, 0.99), + ] + + assert silhouette_score(points, [0, 0, 1, 1, 2, 2]) > 0.8 diff --git a/tests/unit/models/test_speaker_count.py b/tests/unit/models/test_speaker_count.py new file mode 100644 index 0000000..d5fa5e6 --- /dev/null +++ b/tests/unit/models/test_speaker_count.py @@ -0,0 +1,41 @@ +from vox.models.speaker_count import estimate_speaker_count + + +class TestEstimateSpeakerCount: + def test_estimate_when_two_distinct_voices_then_two(self): + points = [(1.0, 0.0), (0.97, 0.05), (0.0, 1.0), (0.05, 0.97)] + + assert estimate_speaker_count(points) == 2 + + def test_estimate_when_three_distinct_voices_then_three(self): + points = [ + (1.0, 0.0, 0.0), + (0.97, 0.05, 0.0), + (0.0, 1.0, 0.0), + (0.0, 0.97, 0.05), + (0.0, 0.0, 1.0), + (0.05, 0.0, 0.97), + ] + + assert estimate_speaker_count(points) == 3 + + def test_estimate_when_all_the_same_voice_then_one(self): + points = [(1.0, 0.0), (0.999, 0.01), (0.998, 0.02), (1.0, 0.005)] + + assert estimate_speaker_count(points) == 1 + + def test_estimate_when_single_point_then_one(self): + assert estimate_speaker_count([(1.0, 0.0)]) == 1 + + def test_estimate_when_empty_then_one(self): + assert estimate_speaker_count([]) == 1 + + def test_estimate_when_max_speakers_caps_then_respects_it(self): + points = [ + (1.0, 0.0, 0.0), + (0.0, 1.0, 0.0), + (0.0, 0.0, 1.0), + (0.7, 0.7, 0.0), + ] + + assert estimate_speaker_count(points, max_speakers=2) <= 2 diff --git a/tests/unit/models/test_turn_relabeling.py b/tests/unit/models/test_turn_relabeling.py new file mode 100644 index 0000000..c0e60ae --- /dev/null +++ b/tests/unit/models/test_turn_relabeling.py @@ -0,0 +1,45 @@ +from vox.models.speaker_turn import SpeakerTurn +from vox.models.turn_relabeling import merge_labels, relabel_turns + +_TURNS = ( + SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), + SpeakerTurn(start=2.0, end=4.0, speaker="SPEAKER_01"), + SpeakerTurn(start=4.0, end=6.0, speaker="SPEAKER_02"), +) + + +class TestRelabelTurns: + def test_relabel_when_mapping_given_then_applied(self): + result = relabel_turns(_TURNS, {"SPEAKER_02": "SPEAKER_00"}) + + assert [t.speaker for t in result] == [ + "SPEAKER_00", + "SPEAKER_01", + "SPEAKER_00", + ] + + def test_relabel_when_empty_mapping_then_unchanged(self): + assert relabel_turns(_TURNS, {}) == _TURNS + + def test_relabel_when_called_then_bounds_preserved(self): + result = relabel_turns(_TURNS, {"SPEAKER_02": "SPEAKER_00"}) + + assert (result[2].start, result[2].end) == (4.0, 6.0) + + +class TestMergeLabels: + def test_merge_when_two_labels_share_a_group_then_mapped_together(self): + mapping = merge_labels(["SPEAKER_00", "SPEAKER_01", "SPEAKER_02"], [0, 1, 0]) + + assert mapping["SPEAKER_00"] == mapping["SPEAKER_02"] + assert mapping["SPEAKER_01"] != mapping["SPEAKER_00"] + + def test_merge_when_called_then_labels_are_renumbered_from_zero(self): + mapping = merge_labels(["SPEAKER_03", "SPEAKER_07"], [1, 0]) + + assert set(mapping.values()) == {"SPEAKER_00", "SPEAKER_01"} + + def test_merge_when_all_one_group_then_single_label(self): + mapping = merge_labels(["SPEAKER_00", "SPEAKER_01"], [0, 0]) + + assert set(mapping.values()) == {"SPEAKER_00"} From 63bcf67af0a7a49a4980d0d5a771a2483d582c80 Mon Sep 17 00:00:00 2001 From: Yoann Date: Fri, 31 Jul 2026 17:39:01 +0200 Subject: [PATCH 4/5] fix: bound the speaker-count probe to the longest turns On a real 36-minute recording the over-segmenting pass produced 290 labels, not the handful the design assumed. Embedding all of them and running the cubic clustering once per candidate count took ~45 minutes. The probe now samples the 24 longest turns regardless of how many labels came out, estimates the count from those, and re-runs diarization with it. Cost no longer depends on the label count. Measured on that file: 45 min -> 7.3 min (4.9x real time). The phone conversation, previously collapsed onto a single speaker, is now split across two voices. Still open: the estimate returns 2 speakers where the recording holds 3-4, so distinct people still share a label. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Uz3rnGtYHs6YcBNr6MkY1z --- .../adapters/auto_speaker_count_diarizer.py | 33 ++++---- src/vox/models/turn_sampling.py | 9 +++ .../test_auto_speaker_count_diarizer.py | 79 +++++++++++-------- tests/unit/models/test_turn_sampling.py | 36 +++++++++ 4 files changed, 107 insertions(+), 50 deletions(-) create mode 100644 src/vox/models/turn_sampling.py create mode 100644 tests/unit/models/test_turn_sampling.py diff --git a/src/vox/adapters/auto_speaker_count_diarizer.py b/src/vox/adapters/auto_speaker_count_diarizer.py index a8403d8..d85da6c 100644 --- a/src/vox/adapters/auto_speaker_count_diarizer.py +++ b/src/vox/adapters/auto_speaker_count_diarizer.py @@ -1,14 +1,13 @@ from pathlib import Path -from vox.models.clustering import cluster_embeddings from vox.models.speaker_count import DEFAULT_MAX_SPEAKERS, estimate_speaker_count from vox.models.speaker_turn import SpeakerTurn -from vox.models.turn_relabeling import merge_labels, relabel_turns -from vox.models.turn_selection import longest_turn_per_speaker +from vox.models.turn_sampling import longest_turns from vox.ports.diarizer import Diarizer from vox.ports.voice_print_extractor import VoicePrintExtractor OVERSEGMENTATION_THRESHOLD = 0.05 +SAMPLE_SIZE = 24 class AutoSpeakerCountDiarizer: @@ -17,10 +16,12 @@ def __init__( diarizer: Diarizer, extractor: VoicePrintExtractor, max_speakers: int = DEFAULT_MAX_SPEAKERS, + sample_size: int = SAMPLE_SIZE, ): self._diarizer = diarizer self._extractor = extractor self._max_speakers = max_speakers + self._sample_size = sample_size def diarize( self, @@ -29,23 +30,19 @@ def diarize( ) -> tuple[SpeakerTurn, ...]: if num_speakers: return self._diarizer.diarize(audio_path, num_speakers) - turns = self._diarizer.diarize(audio_path, None) - longest = longest_turn_per_speaker(turns) - if len(longest) < 2: - return turns - return self._merge_over_segmented(turns, longest, audio_path) + probe = self._diarizer.diarize(audio_path, None) + sample = longest_turns(probe, self._sample_size) + if len(sample) < 2: + return probe + count = self._estimate_count(audio_path, sample) + return self._diarizer.diarize(audio_path, count) - def _merge_over_segmented( + def _estimate_count( self, - turns: tuple[SpeakerTurn, ...], - longest: dict[str, SpeakerTurn], audio_path: Path, - ) -> tuple[SpeakerTurn, ...]: - labels = list(longest) + sample: tuple[SpeakerTurn, ...], + ) -> int: embeddings = [ - self._extractor.extract(audio_path, longest[a].start, longest[a].end) - for a in labels + self._extractor.extract(audio_path, t.start, t.end) for t in sample ] - count = estimate_speaker_count(embeddings, self._max_speakers) - groups = cluster_embeddings(embeddings, count) - return relabel_turns(turns, merge_labels(labels, groups)) + return estimate_speaker_count(embeddings, self._max_speakers) diff --git a/src/vox/models/turn_sampling.py b/src/vox/models/turn_sampling.py new file mode 100644 index 0000000..1a67386 --- /dev/null +++ b/src/vox/models/turn_sampling.py @@ -0,0 +1,9 @@ +from vox.models.speaker_turn import SpeakerTurn + + +def longest_turns( + turns: tuple[SpeakerTurn, ...], + limit: int, +) -> tuple[SpeakerTurn, ...]: + ordered = sorted(turns, key=lambda t: t.end - t.start, reverse=True) + return tuple(ordered[:limit]) diff --git a/tests/unit/adapters/test_auto_speaker_count_diarizer.py b/tests/unit/adapters/test_auto_speaker_count_diarizer.py index f02ee86..1278d53 100644 --- a/tests/unit/adapters/test_auto_speaker_count_diarizer.py +++ b/tests/unit/adapters/test_auto_speaker_count_diarizer.py @@ -14,8 +14,16 @@ ) -def _labels(turns) -> list[str]: - return [t.speaker for t in turns] +class CountingDiarizer: + """Records every call and replays a canned probe, then a final pass.""" + + def __init__(self, probe): + self._probe = probe + self.calls: list[int | None] = [] + + def diarize(self, audio_path, num_speakers): + self.calls.append(num_speakers) + return self._probe class TestAutoSpeakerCountDiarizer: @@ -35,25 +43,14 @@ def test_diarize_when_count_forced_then_no_embedding_computed(self): assert extractor.extract_calls == [] - def test_diarize_when_single_label_then_returned_as_is(self): + def test_diarize_when_single_turn_then_probe_returned_as_is(self): turns = (SpeakerTurn(start=0.0, end=10.0, speaker="SPEAKER_00"),) auto = AutoSpeakerCountDiarizer(FakeDiarizer(turns), FakeVoicePrintExtractor()) assert auto.diarize(_AUDIO, None) == turns - def test_diarize_when_labels_are_distinct_voices_then_all_kept(self): - extractor = FakeVoicePrintExtractor( - { - (0.0, 10.0): (1.0, 0.0, 0.0), - (10.0, 20.0): (0.0, 1.0, 0.0), - (20.0, 30.0): (0.0, 0.0, 1.0), - } - ) - auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) - - assert len(set(_labels(auto.diarize(_AUDIO, None)))) == 3 - - def test_diarize_when_two_labels_are_the_same_voice_then_merged(self): + def test_diarize_when_auto_then_reruns_with_the_estimated_count(self): + inner = CountingDiarizer(_THREE_LABELS) extractor = FakeVoicePrintExtractor( { (0.0, 10.0): (1.0, 0.0), @@ -61,15 +58,14 @@ def test_diarize_when_two_labels_are_the_same_voice_then_merged(self): (20.0, 30.0): (0.99, 0.01), } ) - auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + auto = AutoSpeakerCountDiarizer(inner, extractor) - result = _labels(auto.diarize(_AUDIO, None)) + auto.diarize(_AUDIO, None) - assert len(set(result)) == 2 - assert result[0] == result[2] - assert result[0] != result[1] + assert inner.calls == [None, 2] - def test_diarize_when_all_labels_same_voice_then_collapses_to_one(self): + def test_diarize_when_all_one_voice_then_final_pass_asks_for_one(self): + inner = CountingDiarizer(_THREE_LABELS) extractor = FakeVoicePrintExtractor( { (0.0, 10.0): (1.0, 0.0), @@ -77,21 +73,40 @@ def test_diarize_when_all_labels_same_voice_then_collapses_to_one(self): (20.0, 30.0): (0.998, 0.02), } ) - auto = AutoSpeakerCountDiarizer(FakeDiarizer(_THREE_LABELS), extractor) + auto = AutoSpeakerCountDiarizer(inner, extractor) + + auto.diarize(_AUDIO, None) - assert len(set(_labels(auto.diarize(_AUDIO, None)))) == 1 + assert inner.calls == [None, 1] - def test_diarize_when_auto_then_uses_longest_turn_of_each_label(self): + def test_diarize_when_many_labels_then_sample_is_capped(self): + many = tuple( + SpeakerTurn(start=float(i), end=float(i) + 1, speaker=f"SPEAKER_{i:03d}") + for i in range(300) + ) extractor = FakeVoicePrintExtractor() + auto = AutoSpeakerCountDiarizer( + CountingDiarizer(many), extractor, sample_size=24 + ) + + auto.diarize(_AUDIO, None) + + assert len(extractor.extract_calls) == 24 + + def test_diarize_when_auto_then_samples_the_longest_turns(self): turns = ( - SpeakerTurn(start=0.0, end=2.0, speaker="SPEAKER_00"), - SpeakerTurn(start=5.0, end=30.0, speaker="SPEAKER_00"), - SpeakerTurn(start=30.0, end=40.0, speaker="SPEAKER_01"), + SpeakerTurn(start=0.0, end=1.0, speaker="SPEAKER_00"), + SpeakerTurn(start=10.0, end=40.0, speaker="SPEAKER_01"), + SpeakerTurn(start=50.0, end=90.0, speaker="SPEAKER_02"), + ) + extractor = FakeVoicePrintExtractor() + auto = AutoSpeakerCountDiarizer( + CountingDiarizer(turns), extractor, sample_size=2 ) - auto = AutoSpeakerCountDiarizer(FakeDiarizer(turns), extractor) auto.diarize(_AUDIO, None) - spans = [call[1:] for call in extractor.extract_calls] - assert (5.0, 30.0) in spans - assert (0.0, 2.0) not in spans + spans = [c[1:] for c in extractor.extract_calls] + assert (50.0, 90.0) in spans + assert (10.0, 40.0) in spans + assert (0.0, 1.0) not in spans diff --git a/tests/unit/models/test_turn_sampling.py b/tests/unit/models/test_turn_sampling.py new file mode 100644 index 0000000..99a2253 --- /dev/null +++ b/tests/unit/models/test_turn_sampling.py @@ -0,0 +1,36 @@ +from vox.models.speaker_turn import SpeakerTurn +from vox.models.turn_sampling import longest_turns + + +def _turn(start, end, speaker="SPEAKER_00"): + return SpeakerTurn(start=start, end=end, speaker=speaker) + + +class TestLongestTurns: + def test_sample_when_fewer_than_limit_then_all_returned(self): + turns = (_turn(0.0, 1.0), _turn(2.0, 5.0)) + + assert len(longest_turns(turns, limit=10)) == 2 + + def test_sample_when_more_than_limit_then_capped(self): + turns = tuple(_turn(i, i + 1) for i in range(100)) + + assert len(longest_turns(turns, limit=40)) == 40 + + def test_sample_when_capped_then_keeps_the_longest(self): + turns = (_turn(0.0, 1.0), _turn(10.0, 40.0), _turn(50.0, 51.0)) + + sampled = longest_turns(turns, limit=1) + + assert sampled[0].start == 10.0 + + def test_sample_when_called_then_ordered_by_duration_desc(self): + turns = (_turn(0.0, 1.0), _turn(10.0, 40.0), _turn(50.0, 55.0)) + + sampled = longest_turns(turns, limit=3) + + durations = [t.end - t.start for t in sampled] + assert durations == sorted(durations, reverse=True) + + def test_sample_when_empty_then_empty(self): + assert longest_turns((), limit=10) == () From a9c6bfb1db7158ee87741e957b6666e78c5bde4d Mon Sep 17 00:00:00 2001 From: Yoann Date: Fri, 31 Jul 2026 19:24:28 +0200 Subject: [PATCH 5/5] fix: sample one turn per label so quiet speakers reach the probe Sampling the globally longest turns concentrated on whichever speaker talks most: on a file where one voice dominates, all 24 sampled turns belonged to it and the count collapsed to 1. The probe now takes the longest turn of each label first, then keeps the longest 24 of those. Also adds pick_speaker_count, which can accept a finer split whose silhouette stays within a tolerance of the best score. It defaults to 1.0, i.e. the plain maximum, because lowering it is not justified yet: at 0.70 a real 3-speaker recording did report 3, but the third label held 2 turns out of 609 while the two genuinely distinct people stayed merged, and 2-speaker files drifted to 3-4. Right number, wrong reason. Verified: every 2-speaker file now returns 2, including the one that previously returned 1. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Uz3rnGtYHs6YcBNr6MkY1z --- .../adapters/auto_speaker_count_diarizer.py | 4 ++- src/vox/models/count_selection.py | 19 ++++++++++ src/vox/models/speaker_count.py | 17 +++++---- .../test_auto_speaker_count_diarizer.py | 36 +++++++++++++++++++ tests/unit/models/test_count_selection.py | 32 +++++++++++++++++ 5 files changed, 98 insertions(+), 10 deletions(-) create mode 100644 src/vox/models/count_selection.py create mode 100644 tests/unit/models/test_count_selection.py diff --git a/src/vox/adapters/auto_speaker_count_diarizer.py b/src/vox/adapters/auto_speaker_count_diarizer.py index d85da6c..8efc6a1 100644 --- a/src/vox/adapters/auto_speaker_count_diarizer.py +++ b/src/vox/adapters/auto_speaker_count_diarizer.py @@ -3,6 +3,7 @@ from vox.models.speaker_count import DEFAULT_MAX_SPEAKERS, estimate_speaker_count from vox.models.speaker_turn import SpeakerTurn from vox.models.turn_sampling import longest_turns +from vox.models.turn_selection import longest_turn_per_speaker from vox.ports.diarizer import Diarizer from vox.ports.voice_print_extractor import VoicePrintExtractor @@ -31,7 +32,8 @@ def diarize( if num_speakers: return self._diarizer.diarize(audio_path, num_speakers) probe = self._diarizer.diarize(audio_path, None) - sample = longest_turns(probe, self._sample_size) + one_per_label = tuple(longest_turn_per_speaker(probe).values()) + sample = longest_turns(one_per_label, self._sample_size) if len(sample) < 2: return probe count = self._estimate_count(audio_path, sample) diff --git a/src/vox/models/count_selection.py b/src/vox/models/count_selection.py new file mode 100644 index 0000000..e425a36 --- /dev/null +++ b/src/vox/models/count_selection.py @@ -0,0 +1,19 @@ +# 1.0 keeps the plain best-scoring count. Lower values accept a finer split +# whose score stays within `tolerance` of the best one: it recovers speakers +# the silhouette tends to merge, at the cost of inventing thin extra ones. +# Measured at 0.70 on a real 3-speaker recording: the count became right while +# the third label held 2 turns out of 609, and 2-speaker files drifted to 3-4. +DEFAULT_TOLERANCE = 1.0 + + +def pick_speaker_count( + scores: dict[int, float], + tolerance: float = DEFAULT_TOLERANCE, +) -> int: + if not scores: + return 1 + best = max(scores.values()) + if best <= 0: + return 1 + floor = tolerance * best + return max(count for count, score in scores.items() if score >= floor) diff --git a/src/vox/models/speaker_count.py b/src/vox/models/speaker_count.py index 07a0333..7edc2bb 100644 --- a/src/vox/models/speaker_count.py +++ b/src/vox/models/speaker_count.py @@ -1,4 +1,5 @@ from vox.models.clustering import cluster_embeddings +from vox.models.count_selection import DEFAULT_TOLERANCE, pick_speaker_count from vox.models.silhouette import silhouette_score from vox.models.voice_matching import Embedding, cosine_similarity @@ -9,26 +10,24 @@ def estimate_speaker_count( embeddings: list[Embedding], max_speakers: int = DEFAULT_MAX_SPEAKERS, + tolerance: float = DEFAULT_TOLERANCE, ) -> int: if len(embeddings) < 2: return 1 if _all_one_voice(embeddings): return 1 - return _best_scoring_count(embeddings, max_speakers) + return pick_speaker_count(_score_each_count(embeddings, max_speakers), tolerance) -def _best_scoring_count( +def _score_each_count( embeddings: list[Embedding], max_speakers: int, -) -> int: +) -> dict[int, float]: ceiling = min(max_speakers, len(embeddings)) - scored = [ - (silhouette_score(embeddings, cluster_embeddings(embeddings, n)), n) + return { + n: silhouette_score(embeddings, cluster_embeddings(embeddings, n)) for n in range(2, ceiling + 1) - ] - if not scored: - return 1 - return max(scored)[1] + } def _all_one_voice(embeddings: list[Embedding]) -> bool: diff --git a/tests/unit/adapters/test_auto_speaker_count_diarizer.py b/tests/unit/adapters/test_auto_speaker_count_diarizer.py index 1278d53..d15e69d 100644 --- a/tests/unit/adapters/test_auto_speaker_count_diarizer.py +++ b/tests/unit/adapters/test_auto_speaker_count_diarizer.py @@ -110,3 +110,39 @@ def test_diarize_when_auto_then_samples_the_longest_turns(self): assert (50.0, 90.0) in spans assert (10.0, 40.0) in spans assert (0.0, 1.0) not in spans + + +class TestSampleDiversity: + def test_diarize_when_one_label_dominates_then_sample_covers_other_labels(self): + # SPEAKER_00 owns the 5 longest turns; SPEAKER_01 speaks briefly once. + turns = ( + *( + SpeakerTurn(start=i * 100.0, end=i * 100.0 + 60, speaker="SPEAKER_00") + for i in range(5) + ), + SpeakerTurn(start=900.0, end=905.0, speaker="SPEAKER_01"), + ) + extractor = FakeVoicePrintExtractor() + auto = AutoSpeakerCountDiarizer( + CountingDiarizer(turns), extractor, sample_size=3 + ) + + auto.diarize(_AUDIO, None) + + spans = [c[1:] for c in extractor.extract_calls] + assert (900.0, 905.0) in spans, "the quiet speaker must be sampled" + + def test_diarize_when_sampling_then_one_turn_per_label(self): + turns = ( + *( + SpeakerTurn(start=i * 100.0, end=i * 100.0 + 60, speaker="SPEAKER_00") + for i in range(5) + ), + SpeakerTurn(start=900.0, end=960.0, speaker="SPEAKER_01"), + ) + extractor = FakeVoicePrintExtractor() + auto = AutoSpeakerCountDiarizer(CountingDiarizer(turns), extractor) + + auto.diarize(_AUDIO, None) + + assert len(extractor.extract_calls) == 2 diff --git a/tests/unit/models/test_count_selection.py b/tests/unit/models/test_count_selection.py new file mode 100644 index 0000000..0ebfa6b --- /dev/null +++ b/tests/unit/models/test_count_selection.py @@ -0,0 +1,32 @@ +from vox.models.count_selection import pick_speaker_count + +# Silhouette scores actually measured on a 36-minute recording holding +# three speakers: the plain maximum picks 2 and merges two people. +_REAL_RECORDING = {2: 0.6967, 3: 0.5290, 4: 0.3696, 5: 0.3354, 6: 0.3113} + +_CLEAR_PAIR = {2: 0.9500, 3: 0.4000, 4: 0.3000} + + +class TestPickSpeakerCount: + def test_pick_when_close_runner_up_then_prefers_the_finer_split(self): + assert pick_speaker_count(_REAL_RECORDING, tolerance=0.70) == 3 + + def test_pick_when_runner_up_is_far_behind_then_keeps_the_best(self): + assert pick_speaker_count(_CLEAR_PAIR, tolerance=0.70) == 2 + + def test_pick_when_tolerance_is_one_then_behaves_like_plain_maximum(self): + assert pick_speaker_count(_REAL_RECORDING, tolerance=1.0) == 2 + + def test_pick_when_empty_then_one(self): + assert pick_speaker_count({}, tolerance=0.70) == 1 + + def test_pick_when_all_scores_are_zero_then_one(self): + assert pick_speaker_count({2: 0.0, 3: 0.0}, tolerance=0.70) == 1 + + def test_pick_when_scores_are_negative_then_one(self): + assert pick_speaker_count({2: -0.3, 3: -0.5}, tolerance=0.70) == 1 + + def test_pick_when_several_within_tolerance_then_takes_the_largest(self): + scores = {2: 1.0, 3: 0.95, 4: 0.90, 5: 0.10} + + assert pick_speaker_count(scores, tolerance=0.70) == 4