HF's zero-shot pipeline scores each candidate label against the input
via a hypothesis template ("This example is {}."), so the label
itself needs to read as natural language for entailment scoring to
work -- feeding it a raw config identifier like "tool_use_agentic" or
"diff_checking" asks the model to judge "This example is
tool_use_agentic.", not a sentence its NLI training ever saw.
Measured live 2026-09-06 against bart-large-mnli with the raw labels:
4 of 9 test prompts landed on the wrong category, every miss also
scoring low (<=0.28) -- confidence and correctness tracked each other,
but the raw labels weren't giving the model enough to work with.
Two changes:
- New _CATEGORY_DESCRIPTIONS maps each config category to a natural-
language description, used as the actual candidate label; the
winning description maps back to its category id for the return
value. A category missing from the map falls back to its raw
string rather than raising, so a newly-added config category
degrades gracefully instead of crashing.
- multi_label=True: the pipeline's default (single-label) normalizes
every candidate's score to sum to 1, so a genuinely good match still
gets dragged down whenever another category is also plausible.
Scoring independently lets a clear match score high on its own
terms.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VRQXz5SYZYVWscxS1QqF6U
252 lines
9.2 KiB
Python
252 lines
9.2 KiB
Python
"""local_encoder.py: the zero-shot classifier backing classifier.mode: local_encoder.
|
|
|
|
Never loads a real model. transformers/torch are not installed in this
|
|
environment (they are optional, per requirements-encoder.txt), and even
|
|
where they are, a unit test has no business paying multi-second model-load
|
|
time. Every test here stubs the lazy import point itself.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import sys
|
|
import types
|
|
from unittest.mock import MagicMock
|
|
|
|
import pytest
|
|
|
|
import local_encoder
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _clear_pipeline_cache():
|
|
"""The module caches loaded pipelines by (model_id, device); a stub
|
|
installed in one test must not leak into the next."""
|
|
local_encoder._pipeline_cache.clear()
|
|
yield
|
|
local_encoder._pipeline_cache.clear()
|
|
|
|
|
|
def _install_fake_transformers(pipeline_factory):
|
|
"""Install a fake `transformers` module in sys.modules so `from
|
|
transformers import pipeline` inside local_encoder resolves to our
|
|
stub, without requiring the real package."""
|
|
fake_module = types.ModuleType("transformers")
|
|
fake_module.pipeline = pipeline_factory
|
|
sys.modules["transformers"] = fake_module
|
|
|
|
|
|
@pytest.fixture
|
|
def restore_sys_modules():
|
|
had = "transformers" in sys.modules
|
|
original = sys.modules.get("transformers")
|
|
yield
|
|
if had:
|
|
sys.modules["transformers"] = original
|
|
else:
|
|
sys.modules.pop("transformers", None)
|
|
|
|
|
|
def test_missing_dependency_raises_an_actionable_message(monkeypatch, restore_sys_modules):
|
|
monkeypatch.delitem(sys.modules, "transformers", raising=False)
|
|
# Simulate an uninstalled package: importing it raises ImportError.
|
|
import builtins
|
|
|
|
real_import = builtins.__import__
|
|
|
|
def fake_import(name, *args, **kwargs):
|
|
if name == "transformers":
|
|
raise ImportError("No module named 'transformers'")
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
monkeypatch.setattr(builtins, "__import__", fake_import)
|
|
|
|
with pytest.raises(ImportError, match="pip install -r requirements-encoder.txt"):
|
|
local_encoder.classify_zero_shot(
|
|
"do a thing", ["a", "b"], model_id="stub-model", device="cpu"
|
|
)
|
|
|
|
|
|
def test_classify_zero_shot_returns_top_label_and_score(restore_sys_modules):
|
|
# The pipeline is scored under the natural-language DESCRIPTION, never
|
|
# the raw category id -- the fake must "win" on the description to
|
|
# exercise the real mapping-back-to-category-id path.
|
|
fake_pipeline_instance = MagicMock(
|
|
return_value={
|
|
"labels": ["writing new code", "casual conversation or a general question"],
|
|
"scores": [0.87, 0.13],
|
|
}
|
|
)
|
|
_install_fake_transformers(lambda *a, **k: fake_pipeline_instance)
|
|
|
|
label, score = local_encoder.classify_zero_shot(
|
|
"fix this bug", ["coding_general", "general_chat"],
|
|
model_id="stub-model", device="cpu",
|
|
)
|
|
assert label == "coding_general"
|
|
assert score == pytest.approx(0.87)
|
|
|
|
|
|
def test_classify_zero_shot_sends_descriptions_not_raw_category_ids(restore_sys_modules):
|
|
"""HF's zero-shot pipeline scores a label against a hypothesis template
|
|
("This example is {}."), so a raw config identifier like
|
|
"tool_use_agentic" is not a sentence its NLI training ever saw. Measured
|
|
live 2026-09-06: this was suppressing scores across the board, not just
|
|
on wrong answers."""
|
|
seen_kwargs = {}
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
def run(task, **call_kwargs):
|
|
seen_kwargs.update(call_kwargs)
|
|
return {"labels": call_kwargs["candidate_labels"], "scores": [0.9, 0.1]}
|
|
return run
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot(
|
|
"read the config then update the manifest",
|
|
["tool_use_agentic", "diff_checking"],
|
|
model_id="stub-model", device="cpu",
|
|
)
|
|
assert seen_kwargs["candidate_labels"] == [
|
|
"a multi-step task that requires using tools, running commands, "
|
|
"or reading and writing files",
|
|
"reviewing a code diff or comparing code changes",
|
|
]
|
|
|
|
|
|
def test_classify_zero_shot_passes_multi_label_true(restore_sys_modules):
|
|
"""Single-label (the pipeline default) normalizes every candidate's score
|
|
to sum to 1, so a genuinely good match still gets dragged down whenever
|
|
another category is also plausible. multi_label scores each
|
|
independently."""
|
|
seen_kwargs = {}
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
def run(task, **call_kwargs):
|
|
seen_kwargs.update(call_kwargs)
|
|
return {"labels": call_kwargs["candidate_labels"], "scores": [0.9]}
|
|
return run
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot("x", ["coding_general"], model_id="m", device="cpu")
|
|
assert seen_kwargs["multi_label"] is True
|
|
|
|
|
|
def test_classify_zero_shot_falls_back_to_raw_string_for_unmapped_category(restore_sys_modules):
|
|
"""A category not in _CATEGORY_DESCRIPTIONS (e.g. a newly-added one in
|
|
config.yaml) must degrade gracefully to using its own raw string as the
|
|
description, not raise."""
|
|
fake_pipeline_instance = MagicMock(
|
|
return_value={"labels": ["some_future_category"], "scores": [0.6]}
|
|
)
|
|
_install_fake_transformers(lambda *a, **k: fake_pipeline_instance)
|
|
|
|
label, score = local_encoder.classify_zero_shot(
|
|
"x", ["some_future_category"], model_id="m", device="cpu",
|
|
)
|
|
assert label == "some_future_category"
|
|
assert score == pytest.approx(0.6)
|
|
|
|
|
|
def test_pipeline_is_built_once_and_reused(restore_sys_modules):
|
|
"""Loading a model per request would be far slower than the LLM call
|
|
this mode replaces -- the pipeline must be cached, not rebuilt."""
|
|
build_calls = []
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
build_calls.append(1)
|
|
return MagicMock(return_value={"labels": ["a"], "scores": [0.9]})
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="stub-model", device="cpu")
|
|
local_encoder.classify_zero_shot("y", ["a"], model_id="stub-model", device="cpu")
|
|
|
|
assert len(build_calls) == 1, "pipeline was rebuilt instead of reused"
|
|
|
|
|
|
def test_different_model_or_device_gets_its_own_cached_pipeline(restore_sys_modules):
|
|
build_calls = []
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
build_calls.append(kwargs.get("model"))
|
|
return MagicMock(return_value={"labels": ["a"], "scores": [0.9]})
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="model-1", device="cpu")
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="model-2", device="cpu")
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="model-1", device="cpu")
|
|
|
|
assert build_calls == ["model-1", "model-2"], "a config change should get a fresh pipeline"
|
|
|
|
|
|
def test_cpu_device_maps_to_transformers_device_minus_one(restore_sys_modules):
|
|
seen_kwargs = {}
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
seen_kwargs.update(kwargs)
|
|
return MagicMock(return_value={"labels": ["a"], "scores": [0.9]})
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="m", device="cpu")
|
|
assert seen_kwargs["device"] == -1
|
|
|
|
|
|
def test_cuda_device_maps_to_transformers_device_zero(restore_sys_modules):
|
|
seen_kwargs = {}
|
|
|
|
def fake_pipeline_factory(*args, **kwargs):
|
|
seen_kwargs.update(kwargs)
|
|
return MagicMock(return_value={"labels": ["a"], "scores": [0.9]})
|
|
|
|
_install_fake_transformers(fake_pipeline_factory)
|
|
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="m", device="cuda")
|
|
assert seen_kwargs["device"] == 0
|
|
|
|
|
|
def test_ensure_available_forces_the_load(restore_sys_modules):
|
|
"""The eager startup check config-load performs when local_encoder mode
|
|
is selected -- must actually attempt the load, not just check the import."""
|
|
build_calls = []
|
|
_install_fake_transformers(
|
|
lambda *a, **k: build_calls.append(1) or MagicMock()
|
|
)
|
|
|
|
local_encoder.ensure_available("stub-model", "cpu")
|
|
assert build_calls == [1]
|
|
|
|
|
|
def test_ensure_available_propagates_missing_dependency(monkeypatch, restore_sys_modules):
|
|
monkeypatch.delitem(sys.modules, "transformers", raising=False)
|
|
import builtins
|
|
|
|
real_import = builtins.__import__
|
|
|
|
def fake_import(name, *args, **kwargs):
|
|
if name == "transformers":
|
|
raise ImportError("No module named 'transformers'")
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
monkeypatch.setattr(builtins, "__import__", fake_import)
|
|
|
|
with pytest.raises(ImportError, match="pip install -r requirements-encoder.txt"):
|
|
local_encoder.ensure_available("stub-model", "cpu")
|
|
|
|
|
|
def test_a_real_failure_other_than_missing_dependency_propagates_unmodified(
|
|
restore_sys_modules,
|
|
):
|
|
"""Any other failure (bad model id, OOM, ...) is NOT swallowed -- the
|
|
caller treats it identically to a local-LLM parse failure."""
|
|
|
|
def blows_up(*args, **kwargs):
|
|
raise RuntimeError("model checkpoint not found")
|
|
|
|
_install_fake_transformers(blows_up)
|
|
|
|
with pytest.raises(RuntimeError, match="model checkpoint not found"):
|
|
local_encoder.classify_zero_shot("x", ["a"], model_id="bad-model", device="cpu")
|