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
143 changes: 143 additions & 0 deletions tests/test_cli.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
"""Command-line interface: output naming and error reporting."""

import importlib
import shutil
import sys

import numpy as np
import pytest

from whisper_mlx.audio import AudioLoadError, load_audio

cli = importlib.import_module("whisper_mlx.cli")


@pytest.fixture
def run_cli(monkeypatch, tmp_path):
"""Run `vayu <args>` with transcription and writing stubbed out.

Returns (exit_code, names written, inputs passed to transcribe()).
"""

def run(*args, transcribe=None):
written, transcribed = [], []

def fake_transcribe(audio, **kwargs):
transcribed.append(audio)
return {"text": "", "segments": [], "language": "en"}

monkeypatch.setattr(cli, "transcribe", transcribe or fake_transcribe)
monkeypatch.setattr(
cli,
"get_writer",
lambda fmt, out: lambda r, name, **k: written.append(name),
)
monkeypatch.setattr(sys, "argv", ["vayu", *args, "-o", str(tmp_path)])
try:
cli.main()
code = 0
except SystemExit as e:
code = e.code
return code, written, transcribed

return run


def test_each_input_gets_its_own_output_name(run_cli):
code, written, _ = run_cli("talks/a.mp3", "b.wav")

assert code == 0
assert written == ["a", "b"]


def test_output_name_applies_to_a_single_input(run_cli):
code, written, _ = run_cli("a.mp3", "--output-name", "custom")

assert (code, written) == (0, ["custom"])


def test_output_name_with_several_inputs_is_rejected(run_cli, capsys):
code, written, _ = run_cli("a.mp3", "b.mp3", "--output-name", "custom")

assert code == 2
assert written == []
assert "--output-name" in capsys.readouterr().err


def test_inputs_with_the_same_stem_are_rejected(run_cli, capsys):
code, written, _ = run_cli("day1/talk.mp3", "day2/talk.mp3")

assert code == 2
assert written == []
assert "talk" in capsys.readouterr().err


def test_stdin_is_read_inside_error_handling(run_cli, monkeypatch, capsys):
def broken_stdin(**kwargs):
raise AudioLoadError("Failed to load audio: invalid data")

monkeypatch.setattr(cli.audio, "load_audio", broken_stdin)

code, written, _ = run_cli("-")

assert code == 1
assert " - -: AudioLoadError" in capsys.readouterr().out


def test_stdin_output_is_named_content(run_cli, monkeypatch):
monkeypatch.setattr(cli.audio, "load_audio", lambda **k: np.zeros(16000))

code, written, transcribed = run_cli("-")

assert (code, written) == (0, ["content"])
assert isinstance(transcribed[0], np.ndarray)


def test_missing_file_is_reported_and_others_still_run(run_cli, capsys):
def transcribe(audio, **kwargs):
if audio == "missing.wav":
raise FileNotFoundError(f"Audio file not found: {audio}")
return {"text": "", "segments": [], "language": "en"}

code, written, _ = run_cli("missing.wav", "present.wav", transcribe=transcribe)

assert code == 1
assert written == ["present"]
out = capsys.readouterr().out
assert "1 file(s) failed" in out
assert "missing.wav: FileNotFoundError" in out


def test_load_audio_reports_missing_file(tmp_path):
with pytest.raises(FileNotFoundError, match="Audio file not found"):
load_audio(str(tmp_path / "nope.wav"))


def test_load_audio_reports_missing_ffmpeg(tmp_path, monkeypatch):
clip = tmp_path / "clip.wav"
clip.write_bytes(b"not audio")
monkeypatch.setenv("PATH", str(tmp_path)) # no ffmpeg in here

with pytest.raises(AudioLoadError, match="ffmpeg was not found"):
load_audio(str(clip))


@pytest.mark.skipif(shutil.which("ffmpeg") is None, reason="needs ffmpeg")
def test_load_audio_reports_undecodable_file(tmp_path):
clip = tmp_path / "clip.mp3"
clip.write_text("this is not an mp3")

with pytest.raises(AudioLoadError, match="Failed to load audio"):
load_audio(str(clip))


def test_transcribe_checks_the_file_before_loading_the_model(monkeypatch, tmp_path):
module = importlib.import_module("whisper_mlx.transcribe")

def no_model(*args, **kwargs):
raise AssertionError("model should not be loaded")

monkeypatch.setattr(module.ModelHolder, "get_model", no_model)

with pytest.raises(FileNotFoundError):
module.transcribe(str(tmp_path / "missing.mp3"))
2 changes: 2 additions & 0 deletions whisper_mlx/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
N_SAMPLES,
SAMPLE_RATE,
TOKENS_PER_SECOND,
AudioLoadError,
load_audio,
log_mel_spectrogram,
pad_or_trim,
Expand Down Expand Up @@ -91,6 +92,7 @@
"DecodingResult",
# Audio processing
"load_audio",
"AudioLoadError",
"log_mel_spectrogram",
"pad_or_trim",
"SAMPLE_RATE",
Expand Down
21 changes: 18 additions & 3 deletions whisper_mlx/audio.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,10 @@
TOKENS_PER_SECOND = SAMPLE_RATE // N_SAMPLES_PER_TOKEN # 20ms per audio token


class AudioLoadError(RuntimeError):
"""ffmpeg is missing or could not decode the audio."""


def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bool = False) -> mx.array:
"""
Open an audio file and read as mono waveform, resampling as necessary
Expand All @@ -35,7 +39,14 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo

Returns
-------
A NumPy array containing the audio waveform, in float32 dtype.
An mx.array containing the audio waveform, in float32 dtype.

Raises
------
FileNotFoundError
If the audio file does not exist.
AudioLoadError
If ffmpeg is not installed or fails to decode the audio.
"""

# This launches a subprocess to decode audio while down-mixing
Expand All @@ -44,7 +55,7 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo
cmd = ["ffmpeg", "-i", "pipe:0"]
else:
if not os.path.isfile(file):
raise ValueError(f"Audio file not found: {file}")
raise FileNotFoundError(f"Audio file not found: {file}")
file = os.path.realpath(file)
cmd = ["ffmpeg", "-nostdin", "-i", file]

Expand All @@ -60,8 +71,12 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo
# fmt: on
try:
out = run(cmd, capture_output=True, check=True).stdout
except FileNotFoundError as e:
raise AudioLoadError(
"ffmpeg was not found on PATH; install it (e.g. `brew install ffmpeg`)"
) from e
except CalledProcessError as e:
raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e
raise AudioLoadError(f"Failed to load audio: {e.stderr.decode()}") from e

return mx.array(np.frombuffer(out, np.int16)).flatten().astype(mx.float32) / 32768.0

Expand Down
56 changes: 33 additions & 23 deletions whisper_mlx/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,7 @@
import pathlib
import sys
import warnings
from dataclasses import dataclass, field
from subprocess import CalledProcessError
from dataclasses import dataclass
from typing import List

from . import audio
Expand Down Expand Up @@ -268,51 +267,62 @@ def main():

errors: List[TranscriptionError] = []

for audio_obj in args.pop("audio"):
if audio_obj == "-":
# receive the contents from stdin rather than read a file
audio_obj = audio.load_audio(from_stdin=True)
audio_paths: List[str] = args.pop("audio")
if output_name is not None and len(audio_paths) > 1:
parser.error("--output-name can only be used with a single audio input")
# "-" reads from stdin; its output is named "content"
output_names = [
output_name or ("content" if path == "-" else pathlib.Path(path).stem)
for path in audio_paths
]
duplicates = sorted({n for n in output_names if output_names.count(n) > 1})
if duplicates:
parser.error(
f"several inputs would write to the same output name: {', '.join(duplicates)}"
)

output_name = output_name or "content"
else:
output_name = output_name or pathlib.Path(audio_obj).stem
for audio_path, name in zip(audio_paths, output_names):
try:
audio_obj = (
audio.load_audio(from_stdin=True) if audio_path == "-" else audio_path
)
result = transcribe(
audio_obj,
path_or_hf_repo=path_or_hf_repo,
batch_size=batch_size,
**args,
)
writer(result, output_name, **writer_args)
writer(result, name, **writer_args)
except FileNotFoundError as e:
logger.error(f"File not found: {audio_obj}")
errors.append(TranscriptionError(audio_obj, "FileNotFoundError", str(e)))
logger.error(f"File not found: {audio_path}")
errors.append(TranscriptionError(audio_path, "FileNotFoundError", str(e)))
if strict:
sys.exit(1)
except CalledProcessError as e:
# FFmpeg or other subprocess failures
stderr_msg = e.stderr.decode() if e.stderr else str(e)
logger.error(f"Audio processing failed for {audio_obj}: {stderr_msg}")
errors.append(TranscriptionError(audio_obj, "CalledProcessError", stderr_msg))
except audio.AudioLoadError as e:
# ffmpeg missing or unable to decode the input
logger.error(f"Audio processing failed for {audio_path}: {e}")
errors.append(TranscriptionError(audio_path, "AudioLoadError", str(e)))
if strict:
sys.exit(1)
except ValueError as e:
# Input validation errors
logger.error(f"Invalid input for {audio_obj}: {e}")
errors.append(TranscriptionError(audio_obj, "ValueError", str(e)))
logger.error(f"Invalid input for {audio_path}: {e}")
errors.append(TranscriptionError(audio_path, "ValueError", str(e)))
if strict:
sys.exit(1)
except MemoryError as e:
except MemoryError:
# Out of memory - always fatal
logger.error(f"Out of memory processing {audio_obj}. Try reducing --batch-size.")
logger.error(
f"Out of memory processing {audio_path}. Try reducing --batch-size."
)
raise
except KeyboardInterrupt:
logger.info("Interrupted by user")
sys.exit(130)
except Exception as e:
# Catch-all for unexpected errors
logger.exception(f"Unexpected error processing {audio_obj}")
errors.append(TranscriptionError(audio_obj, type(e).__name__, str(e)))
logger.exception(f"Unexpected error processing {audio_path}")
errors.append(TranscriptionError(audio_path, type(e).__name__, str(e)))
if strict:
raise

Expand Down
4 changes: 4 additions & 0 deletions whisper_mlx/transcribe.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# Copyright 漏 2023 Apple Inc.

import os
import sys
import warnings
from typing import Any, Dict, List, Optional, Tuple, Union
Expand Down Expand Up @@ -143,6 +144,9 @@ def transcribe(
if isinstance(audio, str):
if not audio:
raise ValueError("Audio path cannot be empty")
if not os.path.isfile(audio):
# fail before loading (and possibly downloading) the model
raise FileNotFoundError(f"Audio file not found: {audio}")
elif isinstance(audio, (np.ndarray, mx.array)):
if audio.size == 0:
raise ValueError("Audio array cannot be empty")
Expand Down
Loading