diff --git a/headroom/transforms/kompress_compressor.py b/headroom/transforms/kompress_compressor.py index cf9e724c4..99cac5b20 100644 --- a/headroom/transforms/kompress_compressor.py +++ b/headroom/transforms/kompress_compressor.py @@ -18,9 +18,10 @@ import contextlib import gc import hashlib import logging +import os import threading from dataclasses import dataclass -from typing import Any +from typing import Any, Literal from ..config import TransformResult from ..onnx_runtime import create_cpu_session_options, trim_process_heap @@ -31,6 +32,12 @@ logger = logging.getLogger(__name__) # Default HuggingFace model ID HF_MODEL_ID = "chopratejas/kompress-base" +KOMPRESS_BACKEND_ENV = "HEADROOM_KOMPRESS_BACKEND" +KOMPRESS_ONNX_INTRA_THREADS_ENV = "HEADROOM_KOMPRESS_ONNX_INTRA_THREADS" +KOMPRESS_ONNX_INTER_THREADS_ENV = "HEADROOM_KOMPRESS_ONNX_INTER_THREADS" +KOMPRESS_COREML_CACHE_DIR_ENV = "HEADROOM_KOMPRESS_COREML_CACHE_DIR" + +KompressBackend = Literal["auto", "onnx", "onnx_cpu", "onnx_coreml", "pytorch", "pytorch_mps"] # Model cache: model_id -> (model, tokenizer, backend) # Supports multiple models loaded simultaneously. @@ -38,6 +45,48 @@ _kompress_cache: dict[str, tuple[Any, Any, str]] = {} _kompress_lock = threading.Lock() +def _selected_backend() -> KompressBackend: + raw = os.environ.get(KOMPRESS_BACKEND_ENV, "auto").strip().lower().replace("-", "_") + aliases = { + "": "auto", + "cpu": "onnx_cpu", + "coreml": "onnx_coreml", + "mps": "pytorch_mps", + "torch": "pytorch", + "torch_mps": "pytorch_mps", + "onnx": "onnx", + "onnx_cpu": "onnx_cpu", + "onnx_coreml": "onnx_coreml", + "pytorch": "pytorch", + "pytorch_mps": "pytorch_mps", + "auto": "auto", + } + return aliases.get(raw, "auto") # type: ignore[return-value] + + +def _env_int(name: str) -> int | None: + raw = os.environ.get(name) + if raw is None or raw.strip() == "": + return None + try: + value = int(raw) + except ValueError: + logger.warning("%s must be an integer, got %r; ignoring", name, raw) + return None + if value <= 0: + logger.warning("%s must be positive, got %r; ignoring", name, raw) + return None + return value + + +def _onnx_session_options(ort: Any) -> Any: + return create_cpu_session_options( + ort, + intra_op_num_threads=_env_int(KOMPRESS_ONNX_INTRA_THREADS_ENV), + inter_op_num_threads=_env_int(KOMPRESS_ONNX_INTER_THREADS_ENV), + ) + + def _bucket_count(value: int) -> str: """Return a coarse, privacy-preserving size bucket.""" if value <= 0: @@ -220,7 +269,11 @@ class _OnnxModel: return (np.array(scores) > 0.5).tolist() -def _load_kompress_onnx(model_id: str) -> tuple[Any, Any, str]: +def _load_kompress_onnx( + model_id: str, + *, + use_coreml: bool = False, +) -> tuple[Any, Any, str]: """Download ONNX INT8 model from HuggingFace and load with onnxruntime.""" import onnxruntime as ort from transformers import AutoTokenizer @@ -234,17 +287,44 @@ def _load_kompress_onnx(model_id: str) -> tuple[Any, Any, str]: logger.info("Downloading Kompress ONNX model from %s ...", model_id) onnx_path = hf_hub_download(model_id, "onnx/kompress-int8.onnx") + backend = "onnx_coreml" if use_coreml else "onnx" + providers: list[Any] + if use_coreml: + from headroom import paths as _paths + + coreml_cache_dir = os.environ.get(KOMPRESS_COREML_CACHE_DIR_ENV, "").strip() + cache_dir = ( + coreml_cache_dir + if coreml_cache_dir + else str(_paths.workspace_dir() / "cache" / "coreml") + ) + os.makedirs(cache_dir, exist_ok=True) + providers = [ + ( + "CoreMLExecutionProvider", + { + "ModelFormat": "NeuralNetwork", + "MLComputeUnits": "ALL", + "RequireStaticInputShapes": "1", + "ModelCacheDirectory": cache_dir, + }, + ), + "CPUExecutionProvider", + ] + else: + providers = ["CPUExecutionProvider"] + session = ort.InferenceSession( onnx_path, - create_cpu_session_options(ort), - providers=["CPUExecutionProvider"], + _onnx_session_options(ort), + providers=providers, ) model = _OnnxModel(session) tokenizer = AutoTokenizer.from_pretrained("answerdotai/ModernBERT-base") - _kompress_cache[model_id] = (model, tokenizer, "onnx") - logger.info("Kompress ONNX INT8 loaded: %s", model_id) - return model, tokenizer, "onnx" + _kompress_cache[model_id] = (model, tokenizer, backend) + logger.info("Kompress ONNX INT8 loaded: %s backend=%s", model_id, backend) + return model, tokenizer, backend def _load_kompress_pytorch(model_id: str, device: str = "auto") -> tuple[Any, Any, str]: @@ -290,16 +370,38 @@ def _load_kompress_pytorch(model_id: str, device: str = "auto") -> tuple[Any, An def _load_kompress(model_id: str = HF_MODEL_ID, device: str = "auto") -> tuple[Any, Any, str]: """Load Kompress model, returns (model, tokenizer, backend). - Try ONNX first (lightweight), fall back to PyTorch. + The default keeps the historic behavior: try ONNX CPU first + (lightweight), then fall back to PyTorch. Operators can override via + HEADROOM_KOMPRESS_BACKEND: + + - auto: ONNX CPU first, then PyTorch. + - onnx / onnx_cpu: force ONNX CPU. + - onnx_coreml: force ONNX Runtime CoreML provider with CPU fallback. + - pytorch: force PyTorch with the configured device. + - pytorch_mps: force PyTorch on Apple's MPS backend. + Models are cached by model_id — multiple models can coexist. """ if model_id in _kompress_cache: return _kompress_cache[model_id] - # Prefer ONNX (50MB onnxruntime vs 800MB torch) + backend = _selected_backend() + if backend in ("onnx", "onnx_cpu"): + return _load_kompress_onnx(model_id, use_coreml=False) + + if backend == "onnx_coreml": + return _load_kompress_onnx(model_id, use_coreml=True) + + if backend in ("pytorch", "pytorch_mps"): + forced_device = "mps" if backend == "pytorch_mps" else device + return _load_kompress_pytorch(model_id, forced_device) + + # Auto mode: preserve stable default behavior. This avoids changing + # compression quality/perf characteristics for existing installs while + # allowing opt-in MPS/CoreML experiments via HEADROOM_KOMPRESS_BACKEND. if _is_onnx_available(): try: - return _load_kompress_onnx(model_id) + return _load_kompress_onnx(model_id, use_coreml=False) except Exception as e: logger.warning("ONNX load failed for %s, trying PyTorch: %s", model_id, e) diff --git a/tests/test_transforms/test_kompress_compressor.py b/tests/test_transforms/test_kompress_compressor.py index cad5a85d0..f8884abb4 100644 --- a/tests/test_transforms/test_kompress_compressor.py +++ b/tests/test_transforms/test_kompress_compressor.py @@ -8,6 +8,7 @@ Covers: - Transform interface: apply() method """ +from types import SimpleNamespace from unittest.mock import MagicMock, patch # ── Import safety (the whole point of the fix) ───────────────────────── @@ -66,6 +67,103 @@ class TestLazyImports: assert result.savings_percentage == 50.0 +class TestKompressBackendSelection: + def test_selected_backend_aliases(self, monkeypatch) -> None: + import headroom.transforms.kompress_compressor as kmod + + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "mps") + assert kmod._selected_backend() == "pytorch_mps" + + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "coreml") + assert kmod._selected_backend() == "onnx_coreml" + + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "cpu") + assert kmod._selected_backend() == "onnx_cpu" + + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "unknown") + assert kmod._selected_backend() == "auto" + + def test_forced_pytorch_mps_backend_uses_mps_device(self, monkeypatch) -> None: + import headroom.transforms.kompress_compressor as kmod + + calls: list[tuple[str, str]] = [] + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "pytorch_mps") + monkeypatch.setattr(kmod, "_kompress_cache", {}) + monkeypatch.setattr( + kmod, + "_load_kompress_pytorch", + lambda model_id, device: calls.append((model_id, device)) + or ("model", "tokenizer", "pytorch"), + ) + + assert kmod._load_kompress("model-a", device="auto") == ("model", "tokenizer", "pytorch") + assert calls == [("model-a", "mps")] + + def test_forced_coreml_backend_uses_onnx_coreml(self, monkeypatch) -> None: + import headroom.transforms.kompress_compressor as kmod + + calls: list[tuple[str, bool]] = [] + monkeypatch.setenv("HEADROOM_KOMPRESS_BACKEND", "onnx_coreml") + monkeypatch.setattr(kmod, "_kompress_cache", {}) + monkeypatch.setattr( + kmod, + "_load_kompress_onnx", + lambda model_id, *, use_coreml=False: calls.append((model_id, use_coreml)) + or ("model", "tokenizer", "onnx_coreml"), + ) + + assert kmod._load_kompress("model-b") == ("model", "tokenizer", "onnx_coreml") + assert calls == [("model-b", True)] + + def test_auto_backend_preserves_onnx_first(self, monkeypatch) -> None: + import headroom.transforms.kompress_compressor as kmod + + calls: list[str] = [] + monkeypatch.delenv("HEADROOM_KOMPRESS_BACKEND", raising=False) + monkeypatch.setattr(kmod, "_kompress_cache", {}) + monkeypatch.setattr(kmod, "_is_onnx_available", lambda: True) + monkeypatch.setattr(kmod, "_is_pytorch_available", lambda: True) + monkeypatch.setattr( + kmod, + "_load_kompress_onnx", + lambda model_id, *, use_coreml=False: calls.append("onnx") + or ("model", "tokenizer", "onnx"), + ) + monkeypatch.setattr( + kmod, + "_load_kompress_pytorch", + lambda model_id, device: calls.append("pytorch") or ("model", "tokenizer", "pytorch"), + ) + + assert kmod._load_kompress("model-c") == ("model", "tokenizer", "onnx") + assert calls == ["onnx"] + + def test_onnx_session_options_read_thread_caps(self, monkeypatch) -> None: + import headroom.transforms.kompress_compressor as kmod + + created: list[SimpleNamespace] = [] + + class FakeSessionOptions: + def __init__(self) -> None: + self.intra_op_num_threads = None + self.inter_op_num_threads = None + self.enable_cpu_mem_arena = True + self.enable_mem_pattern = True + + fake_ort = SimpleNamespace( + SessionOptions=lambda: created.append(FakeSessionOptions()) or created[-1] + ) + monkeypatch.setenv("HEADROOM_KOMPRESS_ONNX_INTRA_THREADS", "2") + monkeypatch.setenv("HEADROOM_KOMPRESS_ONNX_INTER_THREADS", "1") + + options = kmod._onnx_session_options(fake_ort) + + assert options.intra_op_num_threads == 2 + assert options.inter_op_num_threads == 1 + assert options.enable_cpu_mem_arena is False + assert options.enable_mem_pattern is False + + # ── KompressResult ──────────────────────────────────────────────────────