From b1a3feb1086e12ba5668b1dbe3ed060a59b4c998 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 21:35:13 +0000 Subject: [PATCH] Reject quant values that cannot be applied resolve_model_path only used quant when QUANT_REPOS had a build for the model; otherwise it quietly returned the full-precision repo, so LightningWhisperMLX(model="turbo", quant="4bit") loaded an unquantized model, and any string (e.g. "3bit") was accepted. It now raises ValueError for an unknown quant level, for a model with no quantized build (listing the ones that have one), and for quant combined with a full repo path. README and the LightningWhisperMLX docstring list the models with quantized builds. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA --- README.md | 2 ++ tests/test_utils.py | 50 ++++++++++++++++++++++++++++++++++++++++ whisper_mlx/lightning.py | 3 ++- whisper_mlx/utils.py | 26 +++++++++++++++++---- 4 files changed, 75 insertions(+), 6 deletions(-) create mode 100644 tests/test_utils.py diff --git a/README.md b/README.md index 5f94b7e..37ae9e9 100644 --- a/README.md +++ b/README.md @@ -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 | diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 0000000..13af82a --- /dev/null +++ b/tests/test_utils.py @@ -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") diff --git a/whisper_mlx/lightning.py b/whisper_mlx/lightning.py index b10c99c..9182dcf 100644 --- a/whisper_mlx/lightning.py +++ b/whisper_mlx/lightning.py @@ -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}") diff --git a/whisper_mlx/utils.py b/whisper_mlx/utils.py index 3d19d83..a0b23f8 100644 --- a/whisper_mlx/utils.py +++ b/whisper_mlx/utils.py @@ -60,6 +60,7 @@ def format_timestamp( } # Quantized model repos +QUANT_LEVELS = ("4bit", "8bit") QUANT_REPOS = { "tiny": { "4bit": "mlx-community/whisper-tiny-mlx-4bit", @@ -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 ------- @@ -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: