feat(phase6): opt-in k-means weight palettization (quantize_nbits)
Phase 6 of the modernization plan: add weight palettization to the
Core ML converter as an opt-in knob, so the SD1.5 / SDXL UNet can
ship at 1/2, 1/2.7 or 1/4 of its current size with ANE-friendly
inference.
CoreMLConverter (and the LCM converter) gains a `quantize_nbits`
dropdown: `none` (default — identical to pre-Phase-6 behavior and
filenames, so existing cached .mlpackages still resolve) / `8` / `6` /
`4`. The value is encoded as `_q<bits>` after the attn suffix, so the
unquantized model and the three palettized variants coexist on disk
under distinct cache keys.
Implementation
- core/naming.compose_out_name: accepts `quantize_nbits`, validates
against {none, 8, 6, 4}, appends `_q<bits>` (none = empty).
- converter.convert_unet: after ct.convert + before .save, runs
coremltools.optimize.coreml.palettize_weights with
OpPalettizerConfig(mode="kmeans", nbits=...) when the value is not
"none". Adds a `Palettization took Xs` log line.
- converter.convert / nodes.CoreMLConverter.convert: pipe the new arg
through; the ComfyUI node exposes it as a dropdown with default
"none" so existing workflows are unchanged at load time.
- bench/scripts/convert_sd15.py: QUANT_NBITS env knob; uses the
pure compose_out_name (replaces the inline string formatter).
Test infra
- tests/unit/test_characterization_out_name.py: 6 new tests pinning
the `_q<bits>` suffix contract, the "none" passthrough (backward
compat), the cn + lora + quant combination, and the invalid-value
ValueError. Total Tier 0 now at 94.
- Makefile gains `bench-quant` (runs the matrix script) and
`convert-quant` (converts q8, q6, q4 sequentially).
- bench/scripts/quant_matrix.py (new): loads each variant, runs
REPEATS forward passes with a fixed seed, then computes the
noise_pred PSNR of each quantized variant against the unquantized
baseline. Writes bench/results/quant_matrix_<sha>.{json,md}.
README
- New "Quantization (Phase 6, opt-in)" section: tradeoff table
measured on M2 Pro SD1.5 1x512x512 SPLIT_EINSUM (sizes 1641/822/
617/412 MB; fwd 197/187/183/180 ms; PSNR 53.5 / 40.2 / 27.5 dB),
plus per-chip/RAM recommendations.
Default-path safety
- "none" produces the same out_name as Phase 5 -> existing
v1-5-pruned-emaonly_1x512x512_se_unet.mlmodelc is still picked up
unchanged; the m2 golden image test continues to anchor.
This commit is contained in:
@@ -143,3 +143,55 @@ def test_lora_names_from_params_sorts_by_name():
|
||||
|
||||
def test_lora_names_from_params_empty_list():
|
||||
assert lora_names_from_params([]) == []
|
||||
|
||||
|
||||
# ---------- Phase 6: quantize_nbits suffix ---------------------------------
|
||||
|
||||
|
||||
def test_quantize_nbits_none_appends_nothing():
|
||||
"""'none' is the default and must keep the pre-Phase-6 filename so
|
||||
existing cached .mlpackages still resolve."""
|
||||
out = compose_out_name(
|
||||
ckpt_name="dreamshaper_8.safetensors",
|
||||
batch_size=1, width=512, height=512,
|
||||
controlnet_support=False,
|
||||
attention_implementation="SPLIT_EINSUM",
|
||||
quantize_nbits="none",
|
||||
)
|
||||
assert out == "dreamshaper_8_1x512x512_se"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("nbits,suffix", [("4", "_q4"), ("6", "_q6"), ("8", "_q8")])
|
||||
def test_quantize_nbits_appends_q_suffix(nbits, suffix):
|
||||
out = compose_out_name(
|
||||
ckpt_name="dreamshaper_8.safetensors",
|
||||
batch_size=1, width=512, height=512,
|
||||
controlnet_support=False,
|
||||
attention_implementation="SPLIT_EINSUM",
|
||||
quantize_nbits=nbits,
|
||||
)
|
||||
assert out == f"dreamshaper_8_1x512x512_se{suffix}"
|
||||
|
||||
|
||||
def test_quantize_nbits_with_controlnet_and_lora():
|
||||
out = compose_out_name(
|
||||
ckpt_name="dreamshaper_8.safetensors",
|
||||
batch_size=1, width=512, height=512,
|
||||
controlnet_support=True,
|
||||
attention_implementation="SPLIT_EINSUM",
|
||||
lora_names=["a.safetensors"],
|
||||
quantize_nbits="6",
|
||||
)
|
||||
assert out == "dreamshaper_8_a_1x512x512_cn_se_q6"
|
||||
|
||||
|
||||
def test_quantize_nbits_invalid_raises():
|
||||
import pytest as _pytest
|
||||
with _pytest.raises(ValueError, match="quantize_nbits"):
|
||||
compose_out_name(
|
||||
ckpt_name="x.safetensors",
|
||||
batch_size=1, width=512, height=512,
|
||||
controlnet_support=False,
|
||||
attention_implementation="SPLIT_EINSUM",
|
||||
quantize_nbits="16", # not in {none, 8, 6, 4}
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user