Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
]

Expand Down
41 changes: 41 additions & 0 deletions src/vox/adapters/audio_decoding.py
Original file line number Diff line number Diff line change
@@ -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",
"-",
]
50 changes: 50 additions & 0 deletions src/vox/adapters/auto_speaker_count_diarizer.py
Original file line number Diff line number Diff line change
@@ -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)
2 changes: 2 additions & 0 deletions src/vox/adapters/cli/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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:
Expand Down
65 changes: 65 additions & 0 deletions src/vox/adapters/cli/speakers_cmd.py
Original file line number Diff line number Diff line change
@@ -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
34 changes: 34 additions & 0 deletions src/vox/adapters/cli/transcribe_cmd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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(),
),
)


Expand Down
14 changes: 14 additions & 0 deletions src/vox/adapters/cpu_threads.py
Original file line number Diff line number Diff line change
@@ -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)
25 changes: 25 additions & 0 deletions src/vox/adapters/diarization_models.py
Original file line number Diff line number Diff line change
@@ -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
30 changes: 28 additions & 2 deletions src/vox/adapters/disk_file_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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:
Expand All @@ -54,6 +79,7 @@ def _segment_to_dict(segment) -> dict:
"start": segment.start,
"end": segment.end,
"text": segment.text,
"speaker": segment.speaker,
}


Expand Down
24 changes: 24 additions & 0 deletions src/vox/adapters/json_voice_print_store.py
Original file line number Diff line number Diff line change
@@ -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()
)
Loading
Loading