fix(kompress): reject artifacts that fail at run, and prefetch model files at startup

Three related cold-start/robustness gaps found while debugging a user report of
0.12% savings across 722 requests.

1. Artifact fallback was unreachable for run-time failures.

`_create_onnx_session` tries int8-wo, then fp32, then int8 — but only skipped a
candidate when `InferenceSession(...)` construction threw. The int8 weight-only
artifact carries MatMulNBits with bits=8; ORT's CPU kernel only handles 8-bit
through the prepacked MLAS path, so a build or ISA without an 8-bit SQNBitGemm
kernel falls into ComputeBUnpacked, which asserts nbits_ == 4. That raises on
session.run() AFTER construction succeeded, so the fp32 candidate was never
reached: 207 consecutive per-request failures over three days with ML
compression silently dead. A two-token smoke run inside the existing candidate
loop makes the fallback fire.

2. A broken model cost an inference on every request, forever.

The per-request handler logged a WARNING and passed through, with no latch — 207
identical lines that read as noise rather than "ML compression is dead". Now
latches to passthrough after 3 CONSECUTIVE failures (any success resets) with one
actionable ERROR naming the artifact override.

3. The model download began on the first request, not at startup.

#2001 correctly moved Kompress off the startup path: on RHEL/CentOS 7-family
hosts, entering cached native init before the port binds segfaults in
libarrow/jemalloc with no Python traceback (#1908) — uncatchable. But it left the
~4-minute DOWNLOAD on the first request, with every request in that window
silently uncompressed.

Downloading is separable from loading. `prefetch_kompress_artifacts` resolves the
files over plain huggingface_hub HTTP and never constructs a session or imports
transformers, so startup can prefetch bytes without touching the boundary #1908
crashes on. Native load stays deferred; the #2001 invariant is preserved and
covered by a test asserting no session is constructed.

Real behavior proof: fallback chain verified against the real repo (int8-wo
rejected, fp32 selected); files-only prefetch resolves in 0.08s warm with
InferenceSession patched to raise. The user's exact run-time rejection can't be
reproduced on ORT 1.21.1 arm64 (it fails at construction here, line 115, not
execution, line 442), so it is pinned with a fake ORT session that loads and then
rejects run().

ruff + mypy clean (508 files); 95 kompress/warmup tests pass.
This commit is contained in:
Tejas Chopra 2026-08-03 10:27:48 -07:00
parent f2c48e26c6
commit 3648565d2a
4 changed files with 393 additions and 11 deletions

View file

@ -4019,6 +4019,22 @@ class ContentRouter(Transform):
logger.debug("HTMLExtractor not available (install trafilatura)")
return self._html_extractor
@staticmethod
def _prefetch_kompress_artifacts_async(kompress_config: Any) -> bool:
"""Start a background download of the Kompress model files, if needed.
Files only see ``prefetch_kompress_artifacts`` for why startup must not
build the model. Returns ``True`` when a prefetch is running.
"""
try:
from .kompress_compressor import HF_MODEL_ID, ensure_background_prefetch
model_id = getattr(kompress_config, "model_id", None) or HF_MODEL_ID
return ensure_background_prefetch(str(model_id))
except Exception as e: # pragma: no cover - defensive; never break startup
logger.debug("Kompress artifact prefetch skipped: %s", e)
return False
def eager_load_compressors(self) -> dict[str, str]:
"""Pre-load compressors at startup to avoid first-request latency.
@ -4033,7 +4049,18 @@ class ContentRouter(Transform):
# 1. ML text compressor: Kompress.
#
# Native model initialization stays out of the blocking startup/lifespan
# path. The existing lazy request path loads Kompress on first use.
# path. The existing lazy request path loads Kompress on first use. This is
# load-bearing, NOT laziness: on RHEL/CentOS 7-family hosts entering cached
# Kompress native init before the port binds segfaults in libarrow/jemalloc
# with no Python traceback (#1908, fixed by #2001) — a crash no try/except
# can catch. Do not call `preload()` here.
#
# What we CAN do at startup is prefetch the model FILES. Downloading is
# pure huggingface_hub HTTP — no ONNX session, no transformers import, so it
# never touches the native path that #1908 crashes on. That removes the real
# cold-start cost: previously the ~4-minute download began on the FIRST
# REQUEST, and every request in that window went silently uncompressed
# behind a single "model not ready" warning.
if self.config.enable_kompress:
compressor = self._get_kompress()
if compressor:
@ -4041,8 +4068,10 @@ class ContentRouter(Transform):
status["kompress"] = "enabled"
status["kompress_backend"] = "unknown"
else:
logger.info("Kompress model preload deferred until first request")
status["kompress"] = "deferred"
if self._prefetch_kompress_artifacts_async(getattr(compressor, "config", None)):
status["kompress_artifacts"] = "prefetching"
logger.info("Kompress model preload deferred until first request")
else:
status["kompress"] = "unavailable"

View file

@ -95,6 +95,12 @@ 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"
KOMPRESS_MAX_CONCURRENT_ENV = "HEADROOM_KOMPRESS_MAX_CONCURRENT"
# Consecutive inference failures before Kompress latches to passthrough for the
# rest of the process. 3 rides out a transient error while still catching a model
# that is broken for this install on the first few requests rather than the 200th.
# ponytail: fixed count, not a rate window — a broken artifact fails every call,
# so there is nothing a window would tell us that three strikes doesn't.
_INFERENCE_FAILURE_LATCH = 3
KOMPRESS_EXECUTION_SEMAPHORE_WAIT_MS_ENV = "HEADROOM_KOMPRESS_EXECUTION_TIMEOUT_MS"
KOMPRESS_EXECUTION_SEMAPHORE_WAIT_MS_DEFAULT = 3000
KOMPRESS_BATCH_SIZE_ENV = "HEADROOM_KOMPRESS_BATCH_SIZE"
@ -599,15 +605,42 @@ def _onnx_filename_candidates() -> tuple[str, ...]:
return _DEFAULT_ONNX_FILENAMES
def _smoke_run(session: Any) -> None:
"""Run one tiny forward pass so a broken artifact fails HERE, not per request.
Some ONNX Runtime builds accept a session and then reject it at execution.
The int8 weight-only artifact carries ``MatMulNBits`` with ``bits=8``; ORT's
CPU kernel only handles 8-bit through the prepacked MLAS path, so a build or
ISA without an 8-bit ``SQNBitGemm`` kernel falls into ``ComputeBUnpacked``,
which hard-asserts ``nbits_ == 4``. That raises on ``session.run()`` after
construction succeeded so a load-only check never sees it and the fp32
fallback below is unreachable. Observed in the wild as 207 consecutive
per-request failures over three days with ML compression silently dead.
Two tokens through the real graph, so it costs milliseconds rather than the
seconds the timed canary takes (kernel dispatch is what fails, not compute).
"""
import numpy as np
session.run(
["final_scores"],
{
"input_ids": np.zeros((1, 2), dtype=np.int64),
"attention_mask": np.ones((1, 2), dtype=np.int64),
},
)
def _create_onnx_session(
model_id: str, providers: list[Any], *, allow_download: bool = True
) -> Any:
"""Resolve and load the model's ONNX artifact, trying candidates in order.
A candidate is skipped on download miss (file not in the repo) or on
session-load failure (e.g. the weight-only int8 artifact uses the
MatMulNBits contrib op, which old onnxruntime builds can't run — those
installs fall through to the fp32 artifact instead of losing Kompress).
A candidate is skipped on download miss (file not in the repo), on
session-load failure, or on smoke-run failure (e.g. the weight-only int8
artifact uses the MatMulNBits contrib op, which some onnxruntime builds
accept at load and then reject at execution those installs fall through to
the fp32 artifact instead of losing Kompress). See :func:`_smoke_run`.
When ``allow_download`` is ``False`` candidates are resolved from the local
cache only; if none is cached, :class:`KompressModelNotCached` is raised
@ -632,15 +665,17 @@ def _create_onnx_session(
ort = onnxruntime
try:
return ort.InferenceSession(
session = ort.InferenceSession(
onnx_path,
_onnx_session_options(ort),
providers=providers,
)
_smoke_run(session)
return session
except Exception as exc:
last_err = exc
logger.warning(
"ONNX artifact %r from %s failed to load (%s); trying next candidate",
"ONNX artifact %r from %s is unusable (%s); trying next candidate",
filename,
model_id,
exc,
@ -954,6 +989,70 @@ def ensure_background_download(model_id: str = HF_MODEL_ID, device: str = "auto"
thread.start()
def prefetch_kompress_artifacts(model_id: str = HF_MODEL_ID) -> bool:
"""Download the model's ONNX artifact to the local cache. No native init.
Deliberately weaker than :func:`warm_kompress_model`: it resolves files over
plain huggingface_hub HTTP and never constructs an ``InferenceSession`` or
imports ``transformers``. That distinction is the whole point entering
Kompress *native* init on the proxy's startup path segfaults in
libarrow/jemalloc on RHEL/CentOS 7-family hosts (#1908, fixed by #2001), so
startup may prefetch bytes but must not build the model.
Stops at the first candidate that resolves: the loader tries them in the same
order, so fetching the rest would be wasted bandwidth.
Returns ``True`` if an artifact is now cached locally.
"""
if model_id in _kompress_cache:
return True
for filename in _onnx_filename_candidates():
try:
hf_hub_download_local_first(model_id, filename, allow_network=True)
return True
except Exception as exc:
logger.debug("Kompress prefetch: %r unavailable for %s: %s", filename, model_id, exc)
return False
def ensure_background_prefetch(model_id: str = HF_MODEL_ID) -> bool:
"""Start a one-shot background artifact prefetch. Non-blocking, idempotent.
Returns ``True`` when a prefetch is running or was started, ``False`` when the
model is already cached (nothing to do) or Kompress isn't installed. Shares the
per-model thread registry with :func:`ensure_background_download` so the two
can't race to fetch the same files.
"""
if not is_kompress_available() or model_id in _kompress_cache:
return False
with _download_threads_lock:
if model_id in _kompress_cache:
return False
existing = _download_threads.get(model_id)
if existing is not None and existing.is_alive():
return True
def _run() -> None:
logger.info("Kompress: prefetching model artifacts for %s ...", model_id)
if prefetch_kompress_artifacts(model_id):
logger.info(
"Kompress: artifact prefetch complete for %s; the model loads on "
"first use without a download stall.",
model_id,
)
else:
logger.warning("Kompress: artifact prefetch found no usable file for %s", model_id)
thread = threading.Thread(
target=_run,
name=f"kompress-prefetch-{model_id.replace('/', '-')}",
daemon=True,
)
_download_threads[model_id] = thread
thread.start()
return True
def warm_kompress_model(
model_id: str = HF_MODEL_ID,
device: str = "cpu",
@ -1083,10 +1182,14 @@ class KompressCompressor(Transform):
def __init__(self, config: KompressConfig | None = None):
self.config = config or KompressConfig()
# Set by the preload canary when inference is too slow to be useful;
# compress()/compress_batch() then pass content through untouched.
# Set by the preload canary when inference is too slow to be useful, or by
# the failure latch when inference raises repeatedly; compress()/
# compress_batch() then pass content through untouched.
self._degraded_reason: str | None = None
self._canary_thread: threading.Thread | None = None
# Consecutive inference failures — reset by any success, so a transient
# error can't accumulate toward the latch across a healthy run.
self._inference_failures: int = 0
def preload(self, *, allow_download: bool = True) -> str:
"""Load the backing model/tokenizer and return the selected backend.
@ -1476,6 +1579,9 @@ class KompressCompressor(Transform):
result.tokens_saved,
)
# A real inference landed — clear the strike count so only CONSECUTIVE
# failures can reach the latch.
self._inference_failures = 0
return result
except KompressModelNotCached:
@ -1485,9 +1591,41 @@ class KompressCompressor(Transform):
)
return self._passthrough(content, n_words)
except Exception as e:
logger.warning("Kompress compression failed: %s", e)
self._record_inference_failure(e)
return self._passthrough(content, n_words)
def _record_inference_failure(self, exc: BaseException) -> None:
"""Log a failed inference, and latch to degraded after repeated failures.
A model that fails once may be transient; one that fails every call is
broken for this process and will never recover on its own. Without a latch
that state is a per-request WARNING forever the reported case logged 207
identical lines across three days while every request silently went
uncompressed, which read as noise rather than "ML compression is dead".
Latching converts it into one actionable line plus a `/debug/warmup`
signal, and stops paying for a call that cannot succeed.
"""
self._inference_failures += 1
if self._degraded_reason is not None:
return
if self._inference_failures < _INFERENCE_FAILURE_LATCH:
logger.warning(
"Kompress compression failed (%d/%d before disabling): %s",
self._inference_failures,
_INFERENCE_FAILURE_LATCH,
exc,
)
return
self._degraded_reason = f"{self._inference_failures} consecutive inference failures: {exc}"
logger.error(
"Kompress inference failed %d times consecutively (%s) — ML compression "
"DISABLED for this run; content passes through uncompressed. Pin a working "
"ONNX artifact via %s=onnx/kompress-fp32.onnx, or report the error above.",
self._inference_failures,
exc,
KOMPRESS_ONNX_FILENAME_ENV,
)
def compress_batch(
self,
contents: list[str],

View file

@ -535,3 +535,146 @@ def test_canary_probe_error_never_breaks_preload(monkeypatch):
assert compressor.preload() == "onnx"
_join_canary(compressor)
assert compressor._degraded_reason is None
# ── Artifact selection: reject at LOAD what would fail at RUN ──────────────────
# Reported case: the int8 weight-only artifact carries MatMulNBits with bits=8.
# ORT's CPU kernel only handles 8-bit via the prepacked MLAS path, so a build
# without an 8-bit SQNBitGemm kernel falls into ComputeBUnpacked, which asserts
# nbits_ == 4. That raises on session.run() AFTER construction succeeded, so the
# load-only candidate loop never saw it and the fp32 fallback was unreachable:
# 207 consecutive per-request failures over three days, ML compression silently
# dead the whole time.
class _FakeOrtSession:
"""Constructs fine; optionally rejects execution the way ORT's CPU kernel does."""
def __init__(self, path: str, *, fails_at_run: bool):
self.path = path
self._fails_at_run = fails_at_run
self.runs = 0
def run(self, outputs, feeds):
self.runs += 1
if self._fails_at_run:
raise RuntimeError(
"[ONNXRuntimeError] : 6 : RUNTIME_EXCEPTION : Non-zero status code "
"returned while running MatMulNBits node ... nbits_ == 4 was false. "
"Only 4b quantization is supported for unpacked compute."
)
import numpy as np
return [np.zeros((1, 2), dtype=np.float32)]
def _install_fake_ort(monkeypatch, *, run_fails_for: set[str]):
"""Patch onnxruntime so InferenceSession succeeds but run() may not."""
created: list[_FakeOrtSession] = []
class _FakeOrt:
@staticmethod
def SessionOptions(): # noqa: N802 - mirrors the ORT API
return object()
@staticmethod
def InferenceSession(path, options=None, providers=None): # noqa: N802
session = _FakeOrtSession(path, fails_at_run=any(bad in path for bad in run_fails_for))
created.append(session)
return session
monkeypatch.setitem(__import__("sys").modules, "onnxruntime", _FakeOrt)
monkeypatch.setattr(kc, "_onnx_session_options", lambda _ort: object())
monkeypatch.setattr(kc, "hf_hub_download_local_first", lambda repo, fn, **kw: f"/cache/{fn}")
return created
def test_run_time_artifact_rejection_falls_through_to_next_candidate(monkeypatch, caplog):
"""A session that loads then fails at run must be skipped, not returned."""
created = _install_fake_ort(monkeypatch, run_fails_for={"int8-wo"})
with caplog.at_level("WARNING"):
session = kc._create_onnx_session("org/model", ["CPUExecutionProvider"])
# int8-wo was constructed, smoke-run, rejected; fp32 was selected instead.
assert "int8-wo" in created[0].path
assert created[0].runs == 1
assert "kompress-fp32.onnx" in session.path
assert "unusable" in caplog.text
def test_healthy_artifact_is_selected_after_one_smoke_run(monkeypatch):
created = _install_fake_ort(monkeypatch, run_fails_for=set())
session = kc._create_onnx_session("org/model", ["CPUExecutionProvider"])
# First candidate works, so no fallback and exactly one probe.
assert session is created[0]
assert len(created) == 1
assert session.runs == 1
def test_all_artifacts_failing_at_run_raises_rather_than_returning_a_dead_session(monkeypatch):
_install_fake_ort(monkeypatch, run_fails_for={"onnx/"})
with pytest.raises(FileNotFoundError, match="No loadable ONNX artifact"):
kc._create_onnx_session("org/model", ["CPUExecutionProvider"])
# ── Failure latch: a broken model stops costing us every request ───────────────
def test_repeated_inference_failures_latch_to_passthrough(monkeypatch, caplog):
class AlwaysFailingModel(FakeModel):
def get_keep_mask(self, input_ids, attention_mask):
self._tick()
raise RuntimeError("MatMulNBits nbits_ == 4 was false")
model = AlwaysFailingModel()
compressor = _make_compressor(monkeypatch, model)
monkeypatch.setenv(KOMPRESS_CANARY_THRESHOLD_ENV, "0") # no canary interference
with caplog.at_level("WARNING"):
for _ in range(kc._INFERENCE_FAILURE_LATCH):
assert compressor.compress(CONTENT_40_WORDS).compressed == CONTENT_40_WORDS
assert compressor._degraded_reason is not None
assert "DISABLED" in caplog.text
calls_at_latch = model.calls
# Latched: further calls short-circuit without touching the model again, so a
# broken artifact can't burn inference on every request for three days.
assert compressor.compress(CONTENT_40_WORDS).compressed == CONTENT_40_WORDS
assert model.calls == calls_at_latch
def test_a_success_resets_the_failure_count(monkeypatch):
class FlakyModel(FakeModel):
def __init__(self):
super().__init__()
self.fail_next = True
def get_keep_mask(self, input_ids, attention_mask):
if self.fail_next:
self._tick()
raise RuntimeError("transient")
return super().get_keep_mask(input_ids, attention_mask)
model = FlakyModel()
compressor = _make_compressor(monkeypatch, model)
monkeypatch.setenv(KOMPRESS_CANARY_THRESHOLD_ENV, "0")
# Two failures, then a success, then two more failures: never 3 in a row.
for _ in range(kc._INFERENCE_FAILURE_LATCH - 1):
compressor.compress(CONTENT_40_WORDS)
assert compressor._inference_failures == kc._INFERENCE_FAILURE_LATCH - 1
model.fail_next = False
compressor.compress(CONTENT_40_WORDS)
assert compressor._inference_failures == 0
assert compressor._degraded_reason is None
model.fail_next = True
for _ in range(kc._INFERENCE_FAILURE_LATCH - 1):
compressor.compress(CONTENT_40_WORDS)
assert compressor._degraded_reason is None

View file

@ -127,6 +127,9 @@ def test_eager_load_defers_kompress_regardless_of_cache_state(monkeypatch, cache
router = _router_kompress_only()
stub = _StubCompressor(cached=cache_state == "cached")
monkeypatch.setattr(router, "_get_kompress", lambda: stub)
# Artifact prefetch is files-only, but it still reaches the network — keep it
# out of this assertion so the test stays about the native-preload boundary.
monkeypatch.setattr(router, "_prefetch_kompress_artifacts_async", lambda _cfg: False)
status = router.eager_load_compressors()
@ -164,6 +167,7 @@ def test_non_kompress_warmups_continue_when_kompress_is_deferred(monkeypatch):
router = _router_kompress_only()
stub = _StubCompressor(cached=True)
monkeypatch.setattr(router, "_get_kompress", lambda: stub)
monkeypatch.setattr(router, "_prefetch_kompress_artifacts_async", lambda _cfg: False)
monkeypatch.setattr("headroom.compression.detector._magika_available", lambda: True)
monkeypatch.setattr("headroom.compression.detector._get_magika", lambda: object())
@ -174,6 +178,74 @@ def test_non_kompress_warmups_continue_when_kompress_is_deferred(monkeypatch):
assert stub.preload_calls == []
# ── Startup artifact prefetch (files only, never native init) ──────────────────
# The cold-start cost #2001 left behind: the ~4-minute model download began on the
# FIRST REQUEST, so every request in that window went silently uncompressed.
# Prefetching FILES at startup is safe because it is plain huggingface_hub HTTP —
# it never constructs an InferenceSession, which is the boundary that segfaults in
# libarrow/jemalloc on RHEL/CentOS 7-family hosts (#1908).
def test_eager_load_starts_artifact_prefetch_without_native_preload(monkeypatch):
router = _router_kompress_only()
stub = _StubCompressor(cached=False)
monkeypatch.setattr(router, "_get_kompress", lambda: stub)
prefetch_calls: list[object] = []
monkeypatch.setattr(
router,
"_prefetch_kompress_artifacts_async",
lambda cfg: (prefetch_calls.append(cfg), True)[1],
)
status = router.eager_load_compressors()
assert status["kompress_artifacts"] == "prefetching"
# The #2001 invariant still holds: no native preload on the startup path.
assert status["kompress"] == "deferred"
assert stub.preload_calls == []
assert len(prefetch_calls) == 1
def test_prefetch_never_constructs_a_session_or_imports_transformers(monkeypatch):
"""The safety property that makes startup prefetch legal at all."""
monkeypatch.setattr(kc, "_kompress_cache", {})
requested: list[str] = []
def fake_local_first(repo_id, filename, *, allow_network=True):
requested.append(filename)
return f"/cache/{filename}"
monkeypatch.setattr(kc, "hf_hub_download_local_first", fake_local_first)
def explode(*args, **kwargs):
raise AssertionError("prefetch must not build the model")
monkeypatch.setattr(kc, "_load_kompress", explode)
monkeypatch.setattr(kc, "_load_kompress_onnx", explode)
assert kc.prefetch_kompress_artifacts("org/model") is True
# Stops at the first candidate that resolves — the loader tries the same order.
assert requested == [kc._onnx_filename_candidates()[0]]
def test_prefetch_reports_false_when_no_artifact_resolves(monkeypatch):
monkeypatch.setattr(kc, "_kompress_cache", {})
def always_missing(repo_id, filename, *, allow_network=True):
raise OSError("not found")
monkeypatch.setattr(kc, "hf_hub_download_local_first", always_missing)
assert kc.prefetch_kompress_artifacts("org/model") is False
def test_background_prefetch_is_noop_when_model_already_cached(monkeypatch):
monkeypatch.setattr(kc, "_kompress_cache", {"org/model": object()})
monkeypatch.setattr(kc, "is_kompress_available", lambda: True)
assert kc.ensure_background_prefetch("org/model") is False
@pytest.mark.asyncio
async def test_proxy_startup_does_not_enter_cached_kompress_native_loader(monkeypatch):
pytest.importorskip("httpx")