diff --git a/mlx_lm/models/qwen3_5_full.py b/mlx_lm/models/qwen3_5_full.py --- a/mlx_lm/models/qwen3_5_full.py +++ b/mlx_lm/models/qwen3_5_full.py @@ -459,9 +459,9 @@ class Model(nn.Module): settings.get("group_size"), settings.get("mode", "affine"), ) - if (bits, group, mode) != (4, 64, "affine"): + if bits not in (4, 5) or group != 64 or mode != "affine": raise ValueError( - "This extension only prepares affine/4-bit/group-64 native checkpoints" + "This extension only prepares affine/4-or-5-bit/group-64 native checkpoints" ) original = shapes[f"{path}.weight"] if original[-1] % group: diff --git a/tests/test_qwen3_5_full.py b/tests/test_qwen3_5_full.py --- a/tests/test_qwen3_5_full.py +++ b/tests/test_qwen3_5_full.py @@ -340,14 +340,17 @@ def test_official_nonquantized_convert_and_reload_preserve_all_tensors( load_model(output, lazy=True) -def test_fixed_affine_quantized_roundtrip_preserves_component_tree(reference, tmp_path): +@pytest.mark.parametrize("bits", [4, 5]) +def test_fixed_affine_quantized_roundtrip_preserves_component_tree( + reference, tmp_path, bits +): from mlx_lm.utils import quantize_model, save_model config, _, _, state, _ = reference model = Model(ModelArgs.from_dict(config)) model.load_weights(list(model.sanitize(state).items()), strict=True) original_names = set(dict(tree_flatten(model.parameters()))) - model, quantized_config = quantize_model(model, config, 64, 4, mode="affine") + model, quantized_config = quantize_model(model, config, 64, bits, mode="affine") save_model(tmp_path, model) save_config(quantized_config, tmp_path / "config.json") loaded, saved_config = load_model(tmp_path, lazy=False, strict=True) @@ -355,7 +358,7 @@ def test_fixed_affine_quantized_roundtrip_preserves_component_tree(reference, tm assert set(actual) == set(expected) assert original_names <= set(actual) assert saved_config["quantization"] == { - "bits": 4, + "bits": bits, "group_size": 64, "mode": "affine", }