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/auto_speaker_count_diarizer.py b/src/vox/adapters/auto_speaker_count_diarizer.py new file mode 100644 index 0000000..8efc6a1 --- /dev/null +++ b/src/vox/adapters/auto_speaker_count_diarizer.py @@ -0,0 +1,50 @@ +from pathlib import Path + +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 + +OVERSEGMENTATION_THRESHOLD = 0.05 +SAMPLE_SIZE = 24 + + +class AutoSpeakerCountDiarizer: + def __init__( + self, + 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, + audio_path: Path, + num_speakers: int | None, + ) -> tuple[SpeakerTurn, ...]: + if num_speakers: + return self._diarizer.diarize(audio_path, num_speakers) + probe = self._diarizer.diarize(audio_path, None) + 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) + return self._diarizer.diarize(audio_path, count) + + def _estimate_count( + self, + audio_path: Path, + sample: tuple[SpeakerTurn, ...], + ) -> int: + embeddings = [ + self._extractor.extract(audio_path, t.start, t.end) for t in sample + ] + return estimate_speaker_count(embeddings, self._max_speakers) 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..fa2e84a 100644 --- a/src/vox/adapters/cli/transcribe_cmd.py +++ b/src/vox/adapters/cli/transcribe_cmd.py @@ -3,19 +3,27 @@ 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 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 +52,18 @@ default="local", help="local (MLX, default) | openai (cloud API)", ) +@click.option("--diarize", is_flag=True, help="Identify speakers (who said what)") +@click.option( + "--no-identify", + is_flag=True, + help="Keep SPEAKER_xx labels even when voices are known", +) +@click.option( + "--speakers", + type=int, + default=None, + help="Force the speaker count (detected automatically if omitted)", +) def transcribe( source, language, @@ -59,6 +79,9 @@ def transcribe( no_cookies, browser, backend, + diarize, + no_identify, + speakers, ): source, language, model = _apply_json_overrides( json_payload, source, language, model @@ -80,6 +103,9 @@ def transcribe( no_clean=no_clean, no_download=no_download, dry_run=dry_run, + diarize=diarize, + no_identify=no_identify, + num_speakers=speakers, ) try: response = use_case.execute(request) @@ -121,6 +147,14 @@ def _build_use_case( transcriber=_build_transcriber(backend), file_writer=DiskFileWriter(), progress=ClickProgressReporter(), + diarizer=AutoSpeakerCountDiarizer( + SherpaDiarizer(clustering_threshold=OVERSEGMENTATION_THRESHOLD), + SherpaVoicePrintExtractor(), + ), + 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/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/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/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/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_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_count.py b/src/vox/models/speaker_count.py new file mode 100644 index 0000000..7edc2bb --- /dev/null +++ b/src/vox/models/speaker_count.py @@ -0,0 +1,38 @@ +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 + +DEFAULT_MAX_SPEAKERS = 10 +SAME_VOICE_SIMILARITY = 0.85 + + +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 pick_speaker_count(_score_each_count(embeddings, max_speakers), tolerance) + + +def _score_each_count( + embeddings: list[Embedding], + max_speakers: int, +) -> dict[int, float]: + ceiling = min(max_speakers, len(embeddings)) + return { + n: silhouette_score(embeddings, cluster_embeddings(embeddings, n)) + for n in range(2, ceiling + 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/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_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/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/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..2e23626 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": "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", + "default": false, + "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/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..4aab980 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 + no_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 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) + 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_auto_speaker_count_diarizer.py b/tests/unit/adapters/test_auto_speaker_count_diarizer.py new file mode 100644 index 0000000..d15e69d --- /dev/null +++ b/tests/unit/adapters/test_auto_speaker_count_diarizer.py @@ -0,0 +1,148 @@ +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"), +) + + +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: + 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_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_auto_then_reruns_with_the_estimated_count(self): + inner = CountingDiarizer(_THREE_LABELS) + 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(inner, extractor) + + auto.diarize(_AUDIO, None) + + assert inner.calls == [None, 2] + + 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), + (10.0, 20.0): (0.999, 0.01), + (20.0, 30.0): (0.998, 0.02), + } + ) + auto = AutoSpeakerCountDiarizer(inner, extractor) + + auto.diarize(_AUDIO, None) + + assert inner.calls == [None, 1] + + 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=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.diarize(_AUDIO, None) + + 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 + + +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/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_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_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 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_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_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_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_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_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"} 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) == () 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..12a4883 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, + "no_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_voices_known_then_labels_replaced_without_any_flag(self): + fix = TranscribeFixture() + fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} + + fix.execute(diarize=True) + + written = fix.file_writer.json_written[0][0] + assert written.segments[0].speaker == "Coco" + + def test_execute_when_no_identify_then_labels_kept(self): + fix = TranscribeFixture() + fix.speaker_identifier.mapping = {"SPEAKER_00": "Coco"} + + 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_no_diarize_then_no_identification(self): + fix = TranscribeFixture() + + 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) + + _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" }, ]