ml-explore / ml-explore/mlx-examples
[whisper] Add Flash Attention and batched decoding for up to 10x speedup
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 9k
- Forks
- 1.2k
- PR merge metrics
- No merged PRs in 30d
Description
Summary
Added Flash Attention and batched segment decoding to mlx-whisper, achieving 9.5x speedup on Apple Silicon.
Changes
1. Flash Attention (whisper.py)
- Replaced manual QKV attention with
mx.fast.scaled_dot_product_attention - Conditional path: uses flash attention by default, falls back to standard attention when
word_timestamps=True(needs QK weights for DTW alignment) - Proper mask handling for autoregressive decoding with KV cache
2. Batched decoding (transcribe.py)
- New
batch_sizeparameter intranscribe()(default=1, fully backward-compatible) - Pre-slices audio into fixed 30s chunks, stacks into batch tensor
(N, 3000, n_mels), decodes simultaneously - Per-segment temperature fallback for quality control
batch_size=1produces identical output to current code
Zero new dependencies. No breaking changes.
Benchmarks (M2 8GB, whisper-small, 5 min Russian audio)
| Mode | Time | Realtime Factor | Speedup |
|---|---|---|---|
| Sequential (batch_size=1) | 9.4s | 4.8x RT | 1x |
| Batched (batch_size=12) | 6.6s | 44.8x RT | 9.5x |
For a 15-hour video: ~20 minutes instead of ~3 hours.
Code
Full implementation with benchmarks: https://github.com/ilyasmukiev/mlx-whisper-pr
- Branch
flash-attention-batch: minimal changes (Flash Attention + batch_size parameter only) - Branch
full-batching-vad-diarize: adds optional VAD (Silero) and speaker diarization
Standalone package: https://github.com/ilyasmukiev/mlx-whisper-fast
Notes
- Could not create a PR directly because
gh repo forkreturns HTTP 502 (repo too large?) - Happy to submit a proper PR once the fork works
- The batched path uses fixed-stride chunking (no dynamic seeking), which is a deliberate trade-off for parallelism — same approach as WhisperX and lightning-whisper-mlx
- Related: Discussion #1275 where batching was acknowledged as possible but not implemented
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reviewing whisper.py and transcribe.py, then compare the proposed flash-attention and batched-decoding implementation in the linked mlx-whisper-pr repository. Confirm the fallback behavior, batch_size backward compatibility, identical batch_size=1 output, and the reported benchmarks before defining the final tests and acceptance criteria.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning, performance
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 35/100