Files
6krrt/tests/test_local_encoder.py
adlee-was-taken 9a3926544a feat(local-encoder): swap from NLI cross-encoding to embedding+centroid classification
Replace transformers.pipeline('zero-shot-classification', bart-large-mnli)
with AutoModel/AutoTokenizer + mean-pool + L2-normalize + cosine
similarity + softmax (BAAI/bge-large-en-v1.5). One forward pass for the
input, cost independent of category count.

Preserves the exact public interface (ensure_available, classify_zero_shot)
— dispatcher.py needs zero changes.

Also updates LocalEncoderConfig's default model and class docstring in
config.py, and rewrites test_local_encoder.py with 13 offline tests
mocking both transformers and torch.

This is Phase 1 of Wave 5.1 in plans/token-waste-waves.md — implementation
only. Phase 2 (re-tuning confidence_threshold against real traffic) and
Phase 3 (delta-detector) are not included.
2026-09-18 01:50:37 -04:00

651 lines
23 KiB
Python

"""local_encoder.py: the embedding+centroid 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 both packages in sys.modules.
"""
from __future__ import annotations
import math
import sys
import types
import pytest
import local_encoder
# ── Mock Tensor ────────────────────────────────────────────────────────
class MockTensor:
"""Minimal tensor emulation for testing the embedding math in local_encoder.
Wraps a 2D list of floats (batch, features) internally and supports
the subset of torch operations used by mean-pool, L2-normalize, cosine
similarity, and softmax.
"""
def __init__(self, data):
# Normalise everything to a 2D list of floats.
if isinstance(data, (int, float)):
data = [[float(data)]]
elif isinstance(data, list):
if not data:
data = [[]]
elif not isinstance(data[0], list):
data = [[float(x) for x in data]]
else:
data = [[float(x) for x in row] for row in data]
self._data = data # list[list[float]]
# ── shape ──────────────────────────────────────────────────────
@property
def shape(self):
if not self._data or not self._data[0]:
return (0, 0)
return (len(self._data), len(self._data[0]))
def size(self, dim=None):
s = self.shape
if dim is None:
return s
return s[dim]
# ── identity / layout ──────────────────────────────────────────
def to(self, *_a, **_k):
return self
def float(self):
return self
def unsqueeze(self, dim):
"""Return self (for expand compatibility)."""
return self
def expand(self, *shape):
return self
@property
def T(self):
"""Transpose: (batch, features) -> (features, batch)."""
if not self._data or not self._data[0]:
return MockTensor([])
rows = len(self._data)
cols = len(self._data[0])
transposed = [[self._data[r][c] for r in range(rows)] for c in range(cols)]
return MockTensor(transposed)
def item(self) -> float:
"""Extract the first element as a scalar float."""
return self._data[0][0]
def __int__(self) -> int:
"""Convert to int (for int(torch.argmax(...)) compatibility)."""
return int(self._data[0][0])
def __float__(self) -> float:
"""Convert to float (for float(torch.argmax(...)) compatibility)."""
return self._data[0][0]
def __index__(self) -> int:
"""Support hex(), oct(), bin() and slice indices."""
return int(self._data[0][0])
def __repr__(self):
return f"MockTensor({self._data})"
def __getitem__(self, idx):
"""Support tensor indexing (e.g. probs[top_idx])."""
if isinstance(idx, int):
if len(self._data) == 1 and len(self._data[0]) > 1:
# Row vector: return the element at column idx as a 0D-like tensor
return MockTensor([[self._data[0][idx]]])
# Column vector: return the row at idx
return MockTensor([self._data[idx]])
if isinstance(idx, MockTensor):
v = int(idx)
return self.__getitem__(v)
if isinstance(idx, slice):
return self
return MockTensor([[self._data[0][idx]]])
# ── element-wise arithmetic ────────────────────────────────────
def __mul__(self, other):
if isinstance(other, (int, float)):
return MockTensor([[x * other for x in row] for row in self._data])
if isinstance(other, MockTensor):
# same-shape element-wise multiply
d2 = other._data
return MockTensor(
[[a * b for a, b in zip(row, row2)] for row, row2 in zip(self._data, d2)]
)
return NotImplemented
def __rmul__(self, other):
if isinstance(other, (int, float)):
return MockTensor([[other * x for x in row] for row in self._data])
return NotImplemented
def __truediv__(self, other):
if isinstance(other, (int, float)):
return MockTensor([[x / other for x in row] for row in self._data])
if isinstance(other, MockTensor):
d2 = other._data
return MockTensor(
[[a / b for a, b in zip(row, row2)] for row, row2 in zip(self._data, d2)]
)
return NotImplemented
def __matmul__(self, other):
"""Dot product / matrix multiply.
For the test case: (1, n) @ (n, 1) -> (1, 1) scalar.
Also handles (1, n) @ (n, m) -> (1, m).
"""
if isinstance(other, MockTensor):
a = self._data
b = other._data
# a is (batch, features), b is (features, n_cats) or (features, 1)
if not a or not b:
return MockTensor(0.0)
n_features = len(a[0])
n_b_cols = len(b[0])
result_row = []
for j in range(n_b_cols):
total = sum(a[0][i] * b[i][j] for i in range(n_features))
result_row.append(total)
return MockTensor([result_row])
return NotImplemented
# ── reduction operations ───────────────────────────────────────
def sum(self, dim=None, keepdim=False):
if not self._data or not self._data[0]:
return MockTensor(0.0)
rows = len(self._data)
cols = len(self._data[0])
if dim is None:
total = sum(sum(row) for row in self._data)
if keepdim:
return MockTensor([[total]])
return MockTensor(total)
if dim == 0:
# sum over rows: produce (1, cols)
result = [sum(self._data[r][c] for r in range(rows)) for c in range(cols)]
if keepdim:
return MockTensor([result])
return MockTensor(result)
if dim == 1:
# sum over columns: produce (rows, 1)
result = [sum(row) for row in self._data]
if keepdim:
return MockTensor([[v] for v in result])
return MockTensor(result)
raise ValueError(f"MockTensor.sum: unsupported dim={dim}")
def clamp(self, min=None, max=None):
if min is not None:
builtin_max = __builtins__["max"] if isinstance(__builtins__, dict) else __builtins__.max
return MockTensor(
[[builtin_max(min, x) for x in row] for row in self._data]
)
return self
def norm(self, dim=None, keepdim=False):
"""L2 norm (Frobenius norm if dim is None)."""
if not self._data or not self._data[0]:
return MockTensor(0.0)
if dim == 1:
# per-row norm
norms = [
math.sqrt(sum(x * x for x in row))
for row in self._data
]
if keepdim:
return MockTensor([[v] for v in norms])
return MockTensor(norms)
# dim is None → Frobenius norm
total = math.sqrt(
sum(sum(x * x for x in row) for row in self._data)
)
if keepdim:
return MockTensor([[total]])
return MockTensor(total)
# ── Test helpers ───────────────────────────────────────────────────────
def _install_fake_modules():
"""Install stubs for ``transformers`` and ``torch`` in sys.modules.
The fake ``transformers`` module provides ``AutoModel`` and
``AutoTokenizer``. The default ``AutoModel.from_pretrained()`` returns a
model whose forward pass returns a ``last_hidden_state`` filled with 1.0,
and ``AutoTokenizer.from_pretrained()`` returns a tokenizer that produces
unit ``attention_mask`` tokens.
Individual tests may override ``from_pretrained`` to customise the model
or tokenizer behaviour.
"""
# ── Mock Model ────────────────────────────────────────────────
class MockModel:
"""Acts as transformers.AutoModel.
``from_pretrained()`` returns an instance whose ``__call__``
returns ``Namespace(last_hidden_state=MockTensor(...))``.
"""
_call_count = 0
@classmethod
def from_pretrained(cls, model_id, **kwargs):
return cls()
def to(self, device_str):
return self
def __call__(self, **kwargs):
MockModel._call_count += 1
# Default: return all-1 embeddings so the math is
# deterministic (mean = 1, norm = sqrt(seq_len)).
attn = kwargs.get("attention_mask", MockTensor([[1]]))
seq_len = attn.shape[1] if attn.shape[0] > 0 else 1
batch = attn.shape[0] if attn.shape[0] > 0 else 1
hidden = MockTensor([[1.0] * seq_len for _ in range(batch)])
return types.SimpleNamespace(last_hidden_state=hidden)
# ── Mock Tokenizer ────────────────────────────────────────────
class MockTokenizer:
"""Acts as transformers.AutoTokenizer.
Always returns fixed ``input_ids`` and unit ``attention_mask``,
both as ``MockTensor``.
"""
@classmethod
def from_pretrained(cls, model_id, **kwargs):
return cls()
def __call__(self, text, **_kw):
# Return a single token with attention_mask = 1
return {
"input_ids": MockTensor([[1]]),
"attention_mask": MockTensor([[1]]),
}
# ── Mock torch ─────────────────────────────────────────────────
def _softmax(x, dim=0):
"""Compute softmax over dim 0 using MockTensor values."""
if isinstance(x, MockTensor):
vals = [row[0] for row in x._data] if x._data else []
else:
vals = list(x)
max_v = max(vals)
exps = [math.exp(v - max_v) for v in vals]
total = sum(exps)
probs = [e / total for e in exps]
return MockTensor([[p] for p in probs])
def _argmax(x, dim=None):
if isinstance(x, MockTensor):
vals = [row[0] for row in x._data] if x._data else []
else:
vals = list(x)
idx = vals.index(max(vals))
return MockTensor([idx])
def _make_tensor(data):
if isinstance(data, MockTensor):
return data
if isinstance(data, (int, float)):
return MockTensor([data])
if isinstance(data, (list, tuple)):
return MockTensor(data)
return MockTensor([0.0])
import contextlib
fake_torch = types.ModuleType("torch")
fake_torch.no_grad = contextlib.nullcontext
fake_torch.Tensor = MockTensor
fake_torch.tensor = _make_tensor
fake_torch.sum = lambda x, dim=None: x.sum(dim=dim)
fake_torch.softmax = _softmax
fake_torch.argmax = _argmax
fake_torch.__version__ = "0.0.0-mock"
fake_transformers = types.ModuleType("transformers")
fake_transformers.AutoModel = MockModel
fake_transformers.AutoTokenizer = MockTokenizer
sys.modules["transformers"] = fake_transformers
sys.modules["torch"] = fake_torch
# Reset call counters
MockModel._call_count = 0
MockModel.from_pretrained = classmethod(lambda cls, mid, **kw: cls())
# ── Fixtures ───────────────────────────────────────────────────────────
@pytest.fixture(autouse=True)
def _clear_caches():
"""All caches in local_encoder must be empty between tests."""
local_encoder._model_cache.clear()
local_encoder._desc_cache.clear()
yield
local_encoder._model_cache.clear()
local_encoder._desc_cache.clear()
@pytest.fixture
def restore_sys_modules():
had_t = "transformers" in sys.modules
orig_t = sys.modules.get("transformers")
had_th = "torch" in sys.modules
orig_th = sys.modules.get("torch")
yield
if had_t:
sys.modules["transformers"] = orig_t
else:
sys.modules.pop("transformers", None)
if had_th:
sys.modules["torch"] = orig_th
else:
sys.modules.pop("torch", None)
@pytest.fixture(autouse=True)
def auto_install_fakes(restore_sys_modules):
"""Install fake transformers+torch for every test by default.
Tests that exercise the missing-dependency path must opt out by
clearing sys.modules within the test body.
"""
_install_fake_modules()
yield
# ── Tests ──────────────────────────────────────────────────────────────
def test_missing_dependency_raises_an_actionable_message(monkeypatch):
"""With transformers/torch removed from sys.modules, the lazy import
raises an ImportError with the actionable install message."""
# Remove the stubs
monkeypatch.delitem(sys.modules, "transformers", raising=False)
monkeypatch.delitem(sys.modules, "torch", 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.classify_zero_shot(
"do a thing", ["a", "b"], model_id="stub-model", device="cpu"
)
def test_classify_zero_shot_returns_top_label_and_score():
"""The category whose description embedding is closest to the input
embedding should be returned as the top label."""
# Pre-set description embeddings: cat_a at [1, 0], cat_b at [0, 1]
# (These are L2-normalised vectors.)
local_encoder._desc_cache[("stub-model", "cpu")] = {
"coding_general": MockTensor([[1.0, 0.0]]),
"debugging": MockTensor([[0.0, 1.0]]),
}
# The classify_zero_shot path will encode the input and compare.
# With the default mock model (all-1 hidden state in a single token),
# after mean-pool + L2-normalize the input embedding is [1.0].
# But the description embeddings are 2D [1.0, 0.0] and [0.0, 1.0].
# For the @ matmul to work, the input needs to be 2D too.
#
# We override the model here: the mock model returns a last_hidden_state
# that after mean-pool produces [0.9, 0.1] (biased toward cat_a).
class CustomModel:
_call_count = 0
@classmethod
def from_pretrained(cls, model_id, **kwargs):
return cls()
def to(self, device_str):
return self
def __call__(self, **kwargs):
CustomModel._call_count += 1
# Return a (1, 2) embedding that will pool to [0.9, 0.1]
return types.SimpleNamespace(
last_hidden_state=MockTensor([[0.9, 0.1]])
)
sys.modules["transformers"].AutoModel = CustomModel
label, score = local_encoder.classify_zero_shot(
"fix this bug", ["coding_general", "debugging"],
model_id="stub-model", device="cpu",
)
assert label == "coding_general"
assert 0.0 <= score <= 1.0
def test_classify_zero_shot_sends_descriptions_not_raw_category_ids():
"""Descriptions from _CATEGORY_DESCRIPTIONS are used to build the
embedding cache, not raw category IDs."""
# The cache starts empty. After classify, it should be populated
# with the category IDs as keys.
assert local_encoder._desc_cache.get(("m", "cpu")) is None
local_encoder.classify_zero_shot(
"hello", ["general_chat"], model_id="m", device="cpu",
)
cached = local_encoder._desc_cache.get(("m", "cpu"))
assert cached is not None
assert "general_chat" in cached
def test_classify_zero_shot_falls_back_to_raw_string_for_unmapped_category():
"""A category not in _CATEGORY_DESCRIPTIONS uses the raw string as
its description and does not raise."""
local_encoder.classify_zero_shot(
"new feature", ["some_future_category"],
model_id="m", device="cpu",
)
# Should not raise. The cache must contain the category.
cached = local_encoder._desc_cache.get(("m", "cpu"))
assert cached is not None
assert "some_future_category" in cached
def test_model_is_built_once_and_reused():
"""AutoModel.from_pretrained must be called only once for repeated
classify calls with the same (model_id, device)."""
# Grab the original from_pretrained to count calls
from_pretrained_calls = []
RealAutoModel = sys.modules["transformers"].AutoModel
orig_fp = RealAutoModel.from_pretrained
@classmethod
def counting_fp(cls, model_id, **kwargs):
from_pretrained_calls.append(1)
return cls()
RealAutoModel.from_pretrained = counting_fp
local_encoder.classify_zero_shot("x", ["a"], model_id="same", device="cpu")
local_encoder.classify_zero_shot("y", ["a"], model_id="same", device="cpu")
assert len(from_pretrained_calls) == 1, "model was rebuilt instead of reused"
RealAutoModel.from_pretrained = orig_fp
def test_different_model_or_device_gets_its_own_cached_model():
"""Different (model_id, device) pairs get separate cached
(model, tokenizer)."""
from_pretrained_seen = []
RealAutoModel = sys.modules["transformers"].AutoModel
@classmethod
def recording_fp(cls, model_id, **kwargs):
from_pretrained_seen.append(model_id)
return cls()
RealAutoModel.from_pretrained = recording_fp
local_encoder.classify_zero_shot("x", ["a"], model_id="model-a", device="cpu")
local_encoder.classify_zero_shot("x", ["a"], model_id="model-b", device="cpu")
local_encoder.classify_zero_shot("x", ["a"], model_id="model-a", device="cpu")
assert from_pretrained_seen == ["model-a", "model-b"], (
"a config change should get a fresh model"
)
def test_device_resolves_correctly():
"""verify_that the model is placed on the right device string."""
# We can inspect the model.to() call via a custom wrapper
device_seen = []
class DeviceCheckingModel:
@classmethod
def from_pretrained(cls, model_id, **kwargs):
return cls()
def to(self, device_str):
device_seen.append(device_str)
return self
def __call__(self, **kwargs):
return types.SimpleNamespace(
last_hidden_state=MockTensor([[1.0]])
)
sys.modules["transformers"].AutoModel = DeviceCheckingModel
# CPU
local_encoder.ensure_available("m", "cpu")
assert "cpu" in device_seen, "cpu device should go to 'cpu'"
# CUDA
device_seen.clear()
local_encoder.ensure_available("m", "cuda")
assert "cuda:0" in device_seen, "cuda device should go to 'cuda:0'"
def test_ensure_available_forces_the_load():
"""ensure_available must trigger model loading."""
loaded = []
RealAutoModel = sys.modules["transformers"].AutoModel
@classmethod
def recording_fp(cls, model_id, **kwargs):
loaded.append(model_id)
return cls()
RealAutoModel.from_pretrained = recording_fp
local_encoder.ensure_available("test-model", "cpu")
assert loaded == ["test-model"]
def test_ensure_available_propagates_missing_dependency(monkeypatch):
"""Missing transformers raises ImportError from ensure_available."""
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)
# Restore torch so the code path reaches the transformers 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():
"""A RuntimeError from from_pretrained must pass through unchanged."""
class BlowsUpModel:
@classmethod
def from_pretrained(cls, model_id, **kwargs):
raise RuntimeError("model checkpoint not found")
def to(self, device_str):
return self
sys.modules["transformers"].AutoModel = BlowsUpModel
with pytest.raises(RuntimeError, match="model checkpoint not found"):
local_encoder.classify_zero_shot(
"x", ["a"], model_id="bad-model", device="cpu"
)
def test_confidence_is_in_range_zero_to_one():
"""Confidence from classify_zero_shot must always be in [0, 1]."""
label, score = local_encoder.classify_zero_shot(
"some text", ["a", "b"], model_id="m", device="cpu",
)
assert isinstance(label, str)
assert 0.0 <= score <= 1.0, f"confidence {score} is outside [0, 1]"
def test_desc_cache_is_populated_on_first_call():
"""After the first classify call, the description cache is populated
for the (model_id, device) pair."""
assert local_encoder._desc_cache.get(("m", "cpu")) is None
local_encoder.classify_zero_shot("hello", ["a"], model_id="m", device="cpu")
cached = local_encoder._desc_cache.get(("m", "cpu"))
assert cached is not None
assert "a" in cached
def test_desc_cache_extended_for_new_categories():
"""Calling with a new category extends the existing cache rather than
rebuilding from scratch."""
local_encoder.classify_zero_shot(
"hello", ["coding_general"], model_id="m", device="cpu",
)
cached_before = local_encoder._desc_cache.get(("m", "cpu"))
assert cached_before is not None
assert "coding_general" in cached_before
assert "debugging" not in cached_before
local_encoder.classify_zero_shot(
"fix bug", ["coding_general", "debugging"], model_id="m", device="cpu",
)
cached_after = local_encoder._desc_cache.get(("m", "cpu"))
assert cached_after is not None
assert "coding_general" in cached_after
assert "debugging" in cached_after