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 README.md
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,8 @@ For reduced memory usage, use quantized models:
whisper = LightningWhisperMLX(model="distil-large-v3", quant="4bit")
```

`quant` accepts `"4bit"` or `"8bit"` and is available for `tiny`, `small`, `medium`, `large-v3` and `distil-large-v3`. Other models raise a `ValueError` rather than silently loading full precision.

## Batch Size Recommendations

| Model | Recommended batch_size | Memory Usage |
Expand Down
50 changes: 50 additions & 0 deletions tests/test_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""Model name resolution."""

import pytest

from whisper_mlx import LightningWhisperMLX
from whisper_mlx.utils import QUANT_REPOS, resolve_model_path


def test_names_resolve_to_repos():
assert resolve_model_path("turbo") == "mlx-community/whisper-turbo"
assert resolve_model_path("tiny") == "mlx-community/whisper-tiny-mlx"


def test_repo_paths_pass_through():
assert resolve_model_path("someone/whisper-custom") == "someone/whisper-custom"


def test_unknown_name_is_rejected():
with pytest.raises(ValueError, match="Unknown model"):
resolve_model_path("gigantic")


@pytest.mark.parametrize("model", sorted(QUANT_REPOS))
@pytest.mark.parametrize("quant", ["4bit", "8bit"])
def test_quantized_builds_resolve(model, quant):
assert resolve_model_path(model, quant) == QUANT_REPOS[model][quant]


def test_unknown_quant_is_rejected():
with pytest.raises(ValueError, match="Unknown quant '3bit'"):
resolve_model_path("tiny", "3bit")


def test_quant_for_model_without_quantized_build_is_rejected():
with pytest.raises(ValueError, match="No 4bit build of 'turbo'"):
resolve_model_path("turbo", "4bit")


def test_quant_with_repo_path_is_rejected():
with pytest.raises(ValueError, match="quant only applies to model names"):
resolve_model_path("someone/whisper-custom", "8bit")


def test_empty_quant_means_full_precision():
assert resolve_model_path("tiny", None) == resolve_model_path("tiny", "")


def test_lightning_wrapper_rejects_unavailable_quant():
with pytest.raises(ValueError, match="No 4bit build"):
LightningWhisperMLX(model="turbo", quant="4bit")
3 changes: 2 additions & 1 deletion whisper_mlx/lightning.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,8 @@ def __init__(
Recommended: 12 for distil models, 6 for large models.

quant : str, optional
Quantization level: "4bit" or "8bit". Only supported for some models.
Quantization level: "4bit" or "8bit". Available for tiny, small,
medium, large-v3 and distil-large-v3; raises ValueError otherwise.
"""
if batch_size < 1:
raise ValueError(f"batch_size must be >= 1, got {batch_size}")
Expand Down
26 changes: 21 additions & 5 deletions whisper_mlx/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ def format_timestamp(
}

# Quantized model repos
QUANT_LEVELS = ("4bit", "8bit")
QUANT_REPOS = {
"tiny": {
"4bit": "mlx-community/whisper-tiny-mlx-4bit",
Expand Down Expand Up @@ -95,7 +96,8 @@ def resolve_model_path(
model : str
Model name (e.g., "tiny", "turbo") or HuggingFace repo path
quant : str, optional
Quantization level: "4bit" or "8bit"
Quantization level: "4bit" or "8bit". Only models listed in
QUANT_REPOS have quantized builds.

Returns
-------
Expand All @@ -105,11 +107,25 @@ def resolve_model_path(
Raises
------
ValueError
If model name is unknown
If the model name is unknown, or quant is invalid or not available
for the model
"""
# Check quantized repos first
if quant and model in QUANT_REPOS and quant in QUANT_REPOS[model]:
return QUANT_REPOS[model][quant]
if quant:
if quant not in QUANT_LEVELS:
raise ValueError(
f"Unknown quant {quant!r}; expected one of {list(QUANT_LEVELS)}"
)
if model in QUANT_REPOS and quant in QUANT_REPOS[model]:
return QUANT_REPOS[model][quant]
if "/" in model:
raise ValueError(
"quant only applies to model names; pass the quantized repo "
f"path as model instead of {model!r}"
)
raise ValueError(
f"No {quant} build of {model!r}. Models with quantized builds: "
f"{list(QUANT_REPOS)}"
)

# Check standard repos
if model in MODEL_REPOS:
Expand Down
Loading