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
38 changes: 38 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,44 @@ whisper = LightningWhisperMLX(model="distil-large-v3", quant="4bit")

Higher batch sizes improve throughput but require more memory. Start with the recommended values and adjust based on your hardware.

### Measuring the speed-up

The speed-up from batching depends on the model, batch size, chip and audio. To measure it on your Mac, time sequential against batched decoding on one of your own files:

```bash
python scripts/benchmark.py audio.mp3 --model distil-large-v3 --batch-size 12
```

The script loads the model and warms it up first, then reports the best of `--runs` timings for `batch_size=1` and for the batch size you chose.

## Batched vs Sequential Decoding

With `batch_size=1`, Vayu decodes like OpenAI's Whisper. Each 30-second window starts where the previous segment ended and is conditioned on the text so far.

With `batch_size > 1`, fixed 30-second windows are decoded together. This is much faster, with some trade-offs:

- Every window in a batch is conditioned on the text from *before* the batch, so `condition_on_previous_text` applies between batches, not between windows.
- Windows don't move to follow the speech. Text cut off at a window boundary is kept as a segment that ends at the boundary, so a word spanning two windows can be split.
- `hallucination_silence_threshold` needs window-by-window seeking and is ignored (with a warning).
- `best_of` is not used. Temperature fallback re-decodes only the windows that fail the quality checks.

Word-level timestamps work in both modes.

## Loading Local Models

`load_model` and `--model` accept a local directory with MLX weights (`config.json` plus `weights.safetensors`, `model.safetensors` or `weights.npz`). For safety, local directories are only loaded from:

- the HuggingFace cache (`~/.cache/huggingface/hub`)
- `/usr/local/share/whisper-mlx`
- directories listed in `WHISPER_MLX_MODEL_DIRS` (separated by `:`)

```bash
export WHISPER_MLX_MODEL_DIRS=~/models:/Volumes/External/whisper
vayu audio.mp3 --model ~/models/whisper-large-v3-mlx
```

Anything else fails with `Model path '...' is not within allowed directories`. HuggingFace repo names (`mlx-community/whisper-turbo`) are downloaded to the cache and are not affected.

## API Reference

### transcribe()
Expand Down
80 changes: 80 additions & 0 deletions scripts/benchmark.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
#!/usr/bin/env python3
"""
Time sequential (batch_size=1) against batched decoding on one audio file.

Usage:
python scripts/benchmark.py audio.mp3 --model distil-large-v3 --batch-size 12

The model is loaded and warmed up before timing, and the audio is decoded once
up front, so the numbers cover transcription only. Use a few minutes of real
speech: short clips fit in one window and can't benefit from batching.
"""

import argparse
import platform
import time

import mlx.core as mx

from whisper_mlx import SAMPLE_RATE, load_audio, transcribe
from whisper_mlx.utils import resolve_model_path


def time_transcription(audio: mx.array, repo: str, batch_size: int, language):
start = time.perf_counter()
result = transcribe(
audio,
path_or_hf_repo=repo,
batch_size=batch_size,
language=language,
verbose=None,
)
return time.perf_counter() - start, result


def main():
parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[1])
parser.add_argument("audio", help="Audio file to transcribe")
parser.add_argument(
"--model", default="distil-large-v3", help="Model name or HuggingFace repo"
)
parser.add_argument("--quant", default=None, choices=["4bit", "8bit"])
parser.add_argument("--batch-size", type=int, default=12)
parser.add_argument(
"--language",
default=None,
help="Language code; set it to leave language detection out of the timings",
)
parser.add_argument("--runs", type=int, default=3, help="Timed runs per setting")
args = parser.parse_args()

repo = resolve_model_path(args.model, args.quant)
audio = load_audio(args.audio)
duration = audio.shape[0] / SAMPLE_RATE

print(f"Model: {repo}")
print(f"Audio: {args.audio} ({duration:.1f}s)")
print(f"System: {platform.platform()}, MLX {mx.__version__}")
print()

# Load the model and compile kernels outside the timed runs
time_transcription(audio[: 30 * SAMPLE_RATE], repo, 1, args.language)

best = {}
for batch_size in dict.fromkeys([1, args.batch_size]):
times = [
time_transcription(audio, repo, batch_size, args.language)[0]
for _ in range(args.runs)
]
best[batch_size] = min(times)
print(
f"batch_size={batch_size:<3} best {best[batch_size]:7.2f}s "
f"({duration / best[batch_size]:5.1f}x real time)"
)

if args.batch_size != 1:
print(f"\nSpeed-up: {best[1] / best[args.batch_size]:.2f}x")


if __name__ == "__main__":
main()
Loading