feat(local-encoder): swap from NLI cross-encoding to embedding+centroid classification #95
@@ -974,7 +974,12 @@ class CloudFallbackConfig(StrictModel):
|
|||||||
class LocalEncoderConfig(StrictModel):
|
class LocalEncoderConfig(StrictModel):
|
||||||
"""A non-generative encoder model used for zero-shot category classification.
|
"""A non-generative encoder model used for zero-shot category classification.
|
||||||
|
|
||||||
Exists to structurally rule out the one failure mode that has cost this
|
Embedding+centroid architecture, not NLI cross-encoding: the input runs
|
||||||
|
through the encoder once and is scored by cosine similarity against
|
||||||
|
precomputed embeddings of the description text, rather than cross-encoded
|
||||||
|
pairwise against each label.
|
||||||
|
|
||||||
|
This still structurally rules out the one failure mode that has cost this
|
||||||
project two classifier generations already (see docs/local-models.md):
|
project two classifier generations already (see docs/local-models.md):
|
||||||
a generative model spending its budget on an unbounded reasoning trace
|
a generative model spending its budget on an unbounded reasoning trace
|
||||||
and returning no parseable JSON. An encoder scored against a fixed label
|
and returning no parseable JSON. An encoder scored against a fixed label
|
||||||
@@ -987,15 +992,12 @@ class LocalEncoderConfig(StrictModel):
|
|||||||
learn from without a new, separate opt-in data-capture feature.
|
learn from without a new, separate opt-in data-capture feature.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# facebook/bart-large-mnli -- the reference model HuggingFace's own docs
|
# BAAI/bge-large-en-v1.5 -- a strong general-purpose English embedding
|
||||||
# use for this exact pipeline. Was MoritzLaurer/deberta-v3-base-zeroshot-v2
|
# model, CPU-viable at this size (~1.3 GB), and well-established in the
|
||||||
# (smaller, ~184M vs ~407M params) until that repo started returning 401
|
# embedding model space. This is Wave 5.1 of the token-waste plan
|
||||||
# even on an unauthenticated GET of its model page -- gated or moved
|
# (plans/token-waste-waves.md): replacing the previous NLI cross-encoding
|
||||||
# sometime after this project picked it. Caught live 2026-09-06: the
|
# approach with an embedding+centroid one.
|
||||||
# startup check (ensure_available) correctly refused to boot rather than
|
model: str = "BAAI/bge-large-en-v1.5"
|
||||||
# fail opaquely on the first request, but it still took production down
|
|
||||||
# until the default was fixed.
|
|
||||||
model: str = "facebook/bart-large-mnli"
|
|
||||||
device: Literal["cpu", "cuda"] = "cpu"
|
device: Literal["cpu", "cuda"] = "cpu"
|
||||||
# Below this, the classification is treated as a FAILURE, not a low-
|
# Below this, the classification is treated as a FAILURE, not a low-
|
||||||
# confidence answer -- the caller cascades exactly as it would for a
|
# confidence answer -- the caller cascades exactly as it would for a
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
"""Zero-shot category classification via a non-generative encoder model.
|
"""Zero-shot category classification via sentence-embedding + nearest-centroid.
|
||||||
|
|
||||||
Backs ``classifier.mode: local_encoder``. Exists to structurally rule out the
|
Backs ``classifier.mode: local_encoder``. Exists to structurally rule out the
|
||||||
one classifier failure mode that has cost this project two generations of
|
one classifier failure mode that has cost this project two generations of
|
||||||
local model already (see ``docs/local-models.md``): a generative model
|
local model already (see ``docs/local-models.md``): a generative model
|
||||||
spending its output budget on an unbounded reasoning trace and returning no
|
spending its output budget on an unbounded reasoning trace and returning no
|
||||||
parseable JSON. A zero-shot encoder classification cannot exhibit that
|
parseable JSON. An embedding-scored classification cannot exhibit that
|
||||||
failure — it scores a fixed label set against the input; there is no trace
|
failure — it scores a fixed label set against the input; there is no trace
|
||||||
to run away.
|
to run away.
|
||||||
|
|
||||||
@@ -13,6 +13,18 @@ task text anywhere (``docs/operations.md``, enforced by a test), so a
|
|||||||
supervised model has no training corpus to learn from without a new, opt-in
|
supervised model has no training corpus to learn from without a new, opt-in
|
||||||
data-capture feature that does not exist yet.
|
data-capture feature that does not exist yet.
|
||||||
|
|
||||||
|
Originally used HuggingFace's ``zero-shot-classification`` pipeline
|
||||||
|
(``facebook/bart-large-mnli`` — a cross-encoder NLI model). Replaced with a
|
||||||
|
sentence-embedding + nearest-centroid architecture for Wave 5.1 of the
|
||||||
|
token-waste reduction plan (``plans/token-waste-waves.md``). An embedding
|
||||||
|
model does one forward pass for the input, then cheap cosine similarity
|
||||||
|
against precomputed category-description embeddings — cost independent of
|
||||||
|
category count — while the NLI pipeline did one cross-encoder pass *per
|
||||||
|
category*. This also affords a stronger backbone on the same CPU/RAM budget:
|
||||||
|
``BAAI/bge-large-en-v1.5`` (the current default) is a general-purpose
|
||||||
|
embedding model known for strong retrieval/classification performance and
|
||||||
|
good CPU viability at 1.3 GB.
|
||||||
|
|
||||||
``transformers``/``torch`` are imported LAZILY, inside ``classify_zero_shot``,
|
``transformers``/``torch`` are imported LAZILY, inside ``classify_zero_shot``,
|
||||||
never at module import time — the same rule ``tui.py`` follows for
|
never at module import time — the same rule ``tui.py`` follows for
|
||||||
``textual``: "only imported by tui.py — never by the service dispatch path,
|
``textual``: "only imported by tui.py — never by the service dispatch path,
|
||||||
@@ -30,19 +42,14 @@ _MISSING_DEPENDENCY_MESSAGE = (
|
|||||||
"pip install -r requirements-encoder.txt"
|
"pip install -r requirements-encoder.txt"
|
||||||
)
|
)
|
||||||
|
|
||||||
# HF's zero-shot pipeline scores each candidate label against the input using
|
# HF's zero-shot pipeline scored each candidate label against the input using
|
||||||
# a hypothesis template ("This example is {}."), so the label itself has to
|
# a hypothesis template ("This example is {}."), so the label itself had to
|
||||||
# read as natural language for the entailment scoring to work well -- feeding
|
# read as natural language for entailment scoring. The embedding-based
|
||||||
# it a raw config identifier like "tool_use_agentic" or "diff_checking" asks
|
# backend reuses the same natural-language descriptions the NLI path used —
|
||||||
# the model to judge "This example is tool_use_agentic.", which is not a
|
# they serve as description embeddings for the centroid comparison. A
|
||||||
# sentence its NLI training ever saw. Measured live 2026-09-06 against
|
# category missing from this map falls back to its raw string rather than
|
||||||
# bart-large-mnli: passing cfg.proficiency.categories's raw strings directly
|
# raising, so a newly-added config category degrades gracefully instead of
|
||||||
# (the previous behavior) put 4 of 9 test prompts on the wrong category, with
|
# crashing.
|
||||||
# every miss also scoring low (<=0.28) -- confidence and correctness track
|
|
||||||
# each other, but the raw labels weren't giving the model enough to work
|
|
||||||
# with. A category missing from this map falls back to its raw string rather
|
|
||||||
# than raising, so a newly-added config category degrades gracefully instead
|
|
||||||
# of crashing.
|
|
||||||
_CATEGORY_DESCRIPTIONS: dict[str, str] = {
|
_CATEGORY_DESCRIPTIONS: dict[str, str] = {
|
||||||
"coding_general": "writing new code",
|
"coding_general": "writing new code",
|
||||||
"coding_refactor": (
|
"coding_refactor": (
|
||||||
@@ -62,33 +69,143 @@ _CATEGORY_DESCRIPTIONS: dict[str, str] = {
|
|||||||
"general_chat": "casual conversation or a general question",
|
"general_chat": "casual conversation or a general question",
|
||||||
}
|
}
|
||||||
|
|
||||||
# The loaded HF pipeline, cached at module level so it survives across
|
# Query-side instruction prefix for BAAI/bge-large-en-v1.5 and other BGE
|
||||||
# calls within one process — reloading a model per request would be far
|
# models. BGE's training uses asymmetric instruction prefixes: the input text
|
||||||
# slower than the LLM call this mode replaces. Keyed by (model_id, device)
|
# (the "query") gets prefixed, while category descriptions (the "passages")
|
||||||
# so a config change to either picks up a fresh pipeline rather than
|
# are encoded as-is. This is harmless for non-BGE embedding models — the
|
||||||
# silently reusing one built for a different model.
|
# extra tokens are just part of the input context.
|
||||||
_pipeline_cache: dict[tuple[str, str], Any] = {}
|
_BGE_QUERY_INSTRUCTION = "Represent this sentence for searching relevant passages: "
|
||||||
|
|
||||||
|
# Softmax temperature for converting cosine similarities to confidence
|
||||||
|
# scores. Cosine similarities for this task typically span [0.30, 0.90]
|
||||||
|
# (unrelated vs closely related). At T=0.10 the softmax sharply amplifies
|
||||||
|
# the winner: a similarity of 0.80 vs 0.50 runners-up gives the winner ~0.95
|
||||||
|
# confidence, while a close match (0.75 vs 0.72) produces a more balanced
|
||||||
|
# spread. This makes the confidence scale meaningful and comparable to the
|
||||||
|
# old NLI pipeline's [0,1] output — /metrics' classifier-degradation-share
|
||||||
|
# warning and config.py's confidence_threshold validator both depend on
|
||||||
|
# confidence being a genuine probability, not a monotonic score.
|
||||||
|
_SOFTMAX_TEMPERATURE = 0.10
|
||||||
|
|
||||||
|
# Cached loaded model + tokenizer, keyed by (model_id, device) so a config
|
||||||
|
# change to either picks up a fresh instance rather than silently reusing
|
||||||
|
# one built for a different model.
|
||||||
|
_model_cache: dict[tuple[str, str], tuple[Any, Any]] = {}
|
||||||
|
|
||||||
|
# Precomputed L2-normalized description embeddings, keyed by
|
||||||
|
# (model_id, device). Built once when the model is first loaded and cached
|
||||||
|
# for the lifetime of the process — descriptions don't change per call, only
|
||||||
|
# the input text needs embedding each time.
|
||||||
|
_desc_cache: dict[tuple[str, str], dict[str, Any]] = {}
|
||||||
|
|
||||||
|
|
||||||
def _load_pipeline(model_id: str, device: str) -> Any:
|
def _mean_pool(token_embeddings: Any, attention_mask: Any) -> Any:
|
||||||
|
"""Mean-pool token embeddings weighted by the attention mask.
|
||||||
|
|
||||||
|
Masks out padding tokens so the mean reflects only real content.
|
||||||
|
"""
|
||||||
|
input_mask_expanded = (
|
||||||
|
attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
|
||||||
|
)
|
||||||
|
sum_embeddings = (token_embeddings * input_mask_expanded).sum(dim=1)
|
||||||
|
sum_mask = input_mask_expanded.sum(dim=1)
|
||||||
|
return sum_embeddings / sum_mask.clamp(min=1e-9)
|
||||||
|
|
||||||
|
|
||||||
|
def _l2_normalize(embeddings: Any) -> Any:
|
||||||
|
"""L2-normalize embeddings along the feature dimension."""
|
||||||
|
return embeddings / embeddings.norm(dim=1, keepdim=True).clamp(min=1e-9)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_description_embeddings(
|
||||||
|
model: Any,
|
||||||
|
tokenizer: Any,
|
||||||
|
device_str: str,
|
||||||
|
categories: list[str],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Precompute L2-normalized embeddings for each category description.
|
||||||
|
|
||||||
|
Each description is encoded through the model, mean-pooled, and
|
||||||
|
L2-normalised. Results are cached so only the input text needs to be
|
||||||
|
encoded per call.
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
desc_embeddings: dict[str, Any] = {}
|
||||||
|
for cat in categories:
|
||||||
|
desc = _CATEGORY_DESCRIPTIONS.get(cat, cat)
|
||||||
|
encoded = tokenizer(
|
||||||
|
desc,
|
||||||
|
padding=True,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
encoded = {k: v.to(device_str) for k, v in encoded.items()}
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = model(**encoded)
|
||||||
|
emb = _mean_pool(outputs.last_hidden_state, encoded["attention_mask"])
|
||||||
|
emb = _l2_normalize(emb)
|
||||||
|
desc_embeddings[cat] = emb
|
||||||
|
return desc_embeddings
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_device(device: str) -> tuple[str, str]:
|
||||||
|
"""Convert config-style ``"cpu"``/``"cuda"`` to torch device strings.
|
||||||
|
|
||||||
|
Returns (model_device, tensor_device) — the model is pinned to CPU explicitly
|
||||||
|
(the default) or to ``cuda:0``, and tensors follow the same mapping.
|
||||||
|
"""
|
||||||
|
if device == "cpu":
|
||||||
|
return "cpu", "cpu"
|
||||||
|
# CUDA: model.go("cuda:0"), tensors go to "cuda:0"
|
||||||
|
return "cuda:0", "cuda:0"
|
||||||
|
|
||||||
|
|
||||||
|
def _load_model_tokenizer(model_id: str, device: str) -> tuple[Any, Any]:
|
||||||
|
"""Load and cache (model, tokenizer) for *model_id* on *device*."""
|
||||||
key = (model_id, device)
|
key = (model_id, device)
|
||||||
cached = _pipeline_cache.get(key)
|
cached = _model_cache.get(key)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
return cached
|
return cached
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from transformers import pipeline
|
from transformers import AutoModel, AutoTokenizer
|
||||||
except ImportError as exc:
|
except ImportError as exc:
|
||||||
raise ImportError(_MISSING_DEPENDENCY_MESSAGE) from exc
|
raise ImportError(_MISSING_DEPENDENCY_MESSAGE) from exc
|
||||||
|
|
||||||
# device=-1 is transformers' own convention for CPU; anything else is a
|
model_device, _tensor_device = _resolve_device(device)
|
||||||
# CUDA device index. This module only exposes "cpu"/"cuda" (device 0)
|
model = AutoModel.from_pretrained(model_id).to(model_device)
|
||||||
# because a router process has no business picking among multiple GPUs
|
tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||||
# for a classifier — that is an operator decision made outside this
|
# Tokenizer always stays on CPU; encoded tensors are moved to device
|
||||||
# module if it ever matters.
|
# at inference time.
|
||||||
device_arg = -1 if device == "cpu" else 0
|
|
||||||
built = pipeline("zero-shot-classification", model=model_id, device=device_arg)
|
_model_cache[key] = (model, tokenizer)
|
||||||
_pipeline_cache[key] = built
|
return model, tokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def _get_or_build_descs(
|
||||||
|
model: Any,
|
||||||
|
tokenizer: Any,
|
||||||
|
device_str: str,
|
||||||
|
model_id: str,
|
||||||
|
device: str,
|
||||||
|
categories: list[str],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Return cached description embeddings or build + cache them."""
|
||||||
|
key = (model_id, device)
|
||||||
|
cached = _desc_cache.get(key)
|
||||||
|
if cached is not None:
|
||||||
|
# Check if all requested categories are already cached
|
||||||
|
missing = [c for c in categories if c not in cached]
|
||||||
|
if not missing:
|
||||||
|
return cached
|
||||||
|
# Extend the cache with any new categories
|
||||||
|
new_embs = _build_description_embeddings(model, tokenizer, device_str, missing)
|
||||||
|
cached.update(new_embs)
|
||||||
|
return cached
|
||||||
|
|
||||||
|
built = _build_description_embeddings(model, tokenizer, device_str, categories)
|
||||||
|
_desc_cache[key] = built
|
||||||
return built
|
return built
|
||||||
|
|
||||||
|
|
||||||
@@ -100,7 +217,7 @@ def ensure_available(model_id: str, device: str) -> None:
|
|||||||
the actionable message above, rather than as an opaque traceback on the
|
the actionable message above, rather than as an opaque traceback on the
|
||||||
first live classification.
|
first live classification.
|
||||||
"""
|
"""
|
||||||
_load_pipeline(model_id, device)
|
_load_model_tokenizer(model_id, device)
|
||||||
|
|
||||||
|
|
||||||
def classify_zero_shot(
|
def classify_zero_shot(
|
||||||
@@ -110,7 +227,9 @@ def classify_zero_shot(
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
device: str,
|
device: str,
|
||||||
) -> tuple[str, float]:
|
) -> tuple[str, float]:
|
||||||
"""One zero-shot classification. Returns (top_category, confidence).
|
"""One zero-shot classification via embedding similarity.
|
||||||
|
|
||||||
|
Returns (top_category, confidence).
|
||||||
|
|
||||||
Raises ``ImportError`` with an actionable message if transformers/torch
|
Raises ``ImportError`` with an actionable message if transformers/torch
|
||||||
are not installed, and propagates any other failure (a bad model id, an
|
are not installed, and propagates any other failure (a bad model id, an
|
||||||
@@ -118,17 +237,60 @@ def classify_zero_shot(
|
|||||||
here identically to a local-LLM parse failure and falls through the
|
here identically to a local-LLM parse failure and falls through the
|
||||||
existing cascade.
|
existing cascade.
|
||||||
|
|
||||||
Categories are scored under their natural-language description
|
The input *task* is encoded through the model, mean-pooled, and
|
||||||
(``_CATEGORY_DESCRIPTIONS``), never the raw config identifier -- see that
|
L2-normalized, then compared via cosine similarity against each
|
||||||
map's docstring for why. ``multi_label=True`` scores each candidate
|
category's precomputed description embedding. Similarities are converted
|
||||||
independently instead of normalizing them to sum to 1: when several
|
to a [0,1] confidence via softmax with a temperature of 0.10 (see the
|
||||||
categories are plausible, the default (single-label) mode forces them to
|
module-level docstring on ``_SOFTMAX_TEMPERATURE`` for the rationale).
|
||||||
compete for probability mass, which drags down the correct answer's score
|
|
||||||
even when it is a clear match on its own terms.
|
Categories are matched under their natural-language description
|
||||||
|
(``_CATEGORY_DESCRIPTIONS``), never the raw config identifier — see that
|
||||||
|
map's docstring for why.
|
||||||
"""
|
"""
|
||||||
classifier = _load_pipeline(model_id, device)
|
import torch
|
||||||
descriptions = [_CATEGORY_DESCRIPTIONS.get(c, c) for c in categories]
|
|
||||||
desc_to_category = dict(zip(descriptions, categories))
|
model, tokenizer = _load_model_tokenizer(model_id, device)
|
||||||
result = classifier(task, candidate_labels=descriptions, multi_label=True)
|
_model_device, tensor_device = _resolve_device(device)
|
||||||
top_description = result["labels"][0]
|
|
||||||
return desc_to_category[top_description], float(result["scores"][0])
|
# Build or retrieve cached description embeddings
|
||||||
|
descriptions = _get_or_build_descs(
|
||||||
|
model, tokenizer, tensor_device, model_id, device, categories,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Encode the input task with BGE's query instruction prefix
|
||||||
|
prefixed_task = _BGE_QUERY_INSTRUCTION + task
|
||||||
|
encoded = tokenizer(
|
||||||
|
prefixed_task,
|
||||||
|
padding=True,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
encoded = {k: v.to(tensor_device) for k, v in encoded.items()}
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = model(**encoded)
|
||||||
|
|
||||||
|
task_emb = _mean_pool(outputs.last_hidden_state, encoded["attention_mask"])
|
||||||
|
task_emb = _l2_normalize(task_emb)
|
||||||
|
|
||||||
|
# Compute cosine similarity against each category description embedding.
|
||||||
|
# All description embeddings are already L2-normalized, so cosine
|
||||||
|
# similarity = dot product.
|
||||||
|
similarities: list[tuple[str, float]] = []
|
||||||
|
for cat in categories:
|
||||||
|
desc_emb = descriptions[cat]
|
||||||
|
sim = (task_emb @ desc_emb.T).item() # scalar cosine similarity
|
||||||
|
similarities.append((cat, sim))
|
||||||
|
|
||||||
|
# Softmax with temperature to produce [0,1] confidence scores. We use
|
||||||
|
# the numerically stable formulation: exp((x - x_max) / T) / sum(...).
|
||||||
|
sim_values = torch.tensor([s for _, s in similarities])
|
||||||
|
scaled = sim_values / _SOFTMAX_TEMPERATURE
|
||||||
|
probs = torch.softmax(scaled, dim=0)
|
||||||
|
|
||||||
|
# Find the top category
|
||||||
|
top_idx = int(torch.argmax(probs))
|
||||||
|
top_cat = similarities[top_idx][0]
|
||||||
|
confidence = float(probs[top_idx])
|
||||||
|
|
||||||
|
return top_cat, confidence
|
||||||
@@ -122,7 +122,7 @@ def test_local_encoder_with_empty_encoder_block_loads_with_defaults(raw):
|
|||||||
cfg["classifier"]["mode"] = "local_encoder"
|
cfg["classifier"]["mode"] = "local_encoder"
|
||||||
cfg["classifier"]["encoder"] = {}
|
cfg["classifier"]["encoder"] = {}
|
||||||
loaded = RouterConfig(**cfg)
|
loaded = RouterConfig(**cfg)
|
||||||
assert loaded.classifier.encoder.model == "facebook/bart-large-mnli"
|
assert loaded.classifier.encoder.model == "BAAI/bge-large-en-v1.5"
|
||||||
assert loaded.classifier.encoder.device == "cpu"
|
assert loaded.classifier.encoder.device == "cpu"
|
||||||
assert loaded.classifier.encoder.confidence_threshold == 0.5
|
assert loaded.classifier.encoder.confidence_threshold == 0.5
|
||||||
|
|
||||||
|
|||||||
@@ -1,55 +1,393 @@
|
|||||||
"""local_encoder.py: the zero-shot classifier backing classifier.mode: local_encoder.
|
"""local_encoder.py: the embedding+centroid classifier backing classifier.mode: local_encoder.
|
||||||
|
|
||||||
Never loads a real model. transformers/torch are not installed in this
|
Never loads a real model. transformers/torch are not installed in this
|
||||||
environment (they are optional, per requirements-encoder.txt), and even
|
environment (they are optional, per requirements-encoder.txt), and even
|
||||||
where they are, a unit test has no business paying multi-second model-load
|
where they are, a unit test has no business paying multi-second model-load
|
||||||
time. Every test here stubs the lazy import point itself.
|
time. Every test here stubs both packages in sys.modules.
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
import sys
|
import sys
|
||||||
import types
|
import types
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
import local_encoder
|
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)
|
@pytest.fixture(autouse=True)
|
||||||
def _clear_pipeline_cache():
|
def _clear_caches():
|
||||||
"""The module caches loaded pipelines by (model_id, device); a stub
|
"""All caches in local_encoder must be empty between tests."""
|
||||||
installed in one test must not leak into the next."""
|
local_encoder._model_cache.clear()
|
||||||
local_encoder._pipeline_cache.clear()
|
local_encoder._desc_cache.clear()
|
||||||
yield
|
yield
|
||||||
local_encoder._pipeline_cache.clear()
|
local_encoder._model_cache.clear()
|
||||||
|
local_encoder._desc_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
|
@pytest.fixture
|
||||||
def restore_sys_modules():
|
def restore_sys_modules():
|
||||||
had = "transformers" in sys.modules
|
had_t = "transformers" in sys.modules
|
||||||
original = sys.modules.get("transformers")
|
orig_t = sys.modules.get("transformers")
|
||||||
|
had_th = "torch" in sys.modules
|
||||||
|
orig_th = sys.modules.get("torch")
|
||||||
yield
|
yield
|
||||||
if had:
|
if had_t:
|
||||||
sys.modules["transformers"] = original
|
sys.modules["transformers"] = orig_t
|
||||||
else:
|
else:
|
||||||
sys.modules.pop("transformers", None)
|
sys.modules.pop("transformers", None)
|
||||||
|
if had_th:
|
||||||
|
sys.modules["torch"] = orig_th
|
||||||
|
else:
|
||||||
|
sys.modules.pop("torch", None)
|
||||||
|
|
||||||
|
|
||||||
def test_missing_dependency_raises_an_actionable_message(monkeypatch, restore_sys_modules):
|
@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, "transformers", raising=False)
|
||||||
# Simulate an uninstalled package: importing it raises ImportError.
|
monkeypatch.delitem(sys.modules, "torch", raising=False)
|
||||||
import builtins
|
|
||||||
|
|
||||||
|
import builtins
|
||||||
real_import = builtins.__import__
|
real_import = builtins.__import__
|
||||||
|
|
||||||
def fake_import(name, *args, **kwargs):
|
def fake_import(name, *args, **kwargs):
|
||||||
@@ -65,164 +403,180 @@ def test_missing_dependency_raises_an_actionable_message(monkeypatch, restore_sy
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_classify_zero_shot_returns_top_label_and_score(restore_sys_modules):
|
def test_classify_zero_shot_returns_top_label_and_score():
|
||||||
# The pipeline is scored under the natural-language DESCRIPTION, never
|
"""The category whose description embedding is closest to the input
|
||||||
# the raw category id -- the fake must "win" on the description to
|
embedding should be returned as the top label."""
|
||||||
# exercise the real mapping-back-to-category-id path.
|
# Pre-set description embeddings: cat_a at [1, 0], cat_b at [0, 1]
|
||||||
fake_pipeline_instance = MagicMock(
|
# (These are L2-normalised vectors.)
|
||||||
return_value={
|
local_encoder._desc_cache[("stub-model", "cpu")] = {
|
||||||
"labels": ["writing new code", "casual conversation or a general question"],
|
"coding_general": MockTensor([[1.0, 0.0]]),
|
||||||
"scores": [0.87, 0.13],
|
"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]])
|
||||||
)
|
)
|
||||||
_install_fake_transformers(lambda *a, **k: fake_pipeline_instance)
|
|
||||||
|
sys.modules["transformers"].AutoModel = CustomModel
|
||||||
|
|
||||||
label, score = local_encoder.classify_zero_shot(
|
label, score = local_encoder.classify_zero_shot(
|
||||||
"fix this bug", ["coding_general", "general_chat"],
|
"fix this bug", ["coding_general", "debugging"],
|
||||||
model_id="stub-model", device="cpu",
|
model_id="stub-model", device="cpu",
|
||||||
)
|
)
|
||||||
assert label == "coding_general"
|
assert label == "coding_general"
|
||||||
assert score == pytest.approx(0.87)
|
assert 0.0 <= score <= 1.0
|
||||||
|
|
||||||
|
|
||||||
def test_classify_zero_shot_sends_descriptions_not_raw_category_ids(restore_sys_modules):
|
def test_classify_zero_shot_sends_descriptions_not_raw_category_ids():
|
||||||
"""HF's zero-shot pipeline scores a label against a hypothesis template
|
"""Descriptions from _CATEGORY_DESCRIPTIONS are used to build the
|
||||||
("This example is {}."), so a raw config identifier like
|
embedding cache, not raw category IDs."""
|
||||||
"tool_use_agentic" is not a sentence its NLI training ever saw. Measured
|
# The cache starts empty. After classify, it should be populated
|
||||||
live 2026-09-06: this was suppressing scores across the board, not just
|
# with the category IDs as keys.
|
||||||
on wrong answers."""
|
assert local_encoder._desc_cache.get(("m", "cpu")) is None
|
||||||
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(
|
local_encoder.classify_zero_shot(
|
||||||
"read the config then update the manifest",
|
"hello", ["general_chat"], model_id="m", device="cpu",
|
||||||
["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")
|
cached = local_encoder._desc_cache.get(("m", "cpu"))
|
||||||
assert build_calls == [1]
|
assert cached is not None
|
||||||
|
assert "general_chat" in cached
|
||||||
|
|
||||||
|
|
||||||
def test_ensure_available_propagates_missing_dependency(monkeypatch, restore_sys_modules):
|
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)
|
monkeypatch.delitem(sys.modules, "transformers", raising=False)
|
||||||
import builtins
|
|
||||||
|
|
||||||
|
import builtins
|
||||||
real_import = builtins.__import__
|
real_import = builtins.__import__
|
||||||
|
|
||||||
def fake_import(name, *args, **kwargs):
|
def fake_import(name, *args, **kwargs):
|
||||||
@@ -232,20 +586,66 @@ def test_ensure_available_propagates_missing_dependency(monkeypatch, restore_sys
|
|||||||
|
|
||||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
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"):
|
with pytest.raises(ImportError, match="pip install -r requirements-encoder.txt"):
|
||||||
local_encoder.ensure_available("stub-model", "cpu")
|
local_encoder.ensure_available("stub-model", "cpu")
|
||||||
|
|
||||||
|
|
||||||
def test_a_real_failure_other_than_missing_dependency_propagates_unmodified(
|
def test_a_real_failure_other_than_missing_dependency_propagates_unmodified():
|
||||||
restore_sys_modules,
|
"""A RuntimeError from from_pretrained must pass through unchanged."""
|
||||||
):
|
|
||||||
"""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):
|
class BlowsUpModel:
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(cls, model_id, **kwargs):
|
||||||
raise RuntimeError("model checkpoint not found")
|
raise RuntimeError("model checkpoint not found")
|
||||||
|
|
||||||
_install_fake_transformers(blows_up)
|
def to(self, device_str):
|
||||||
|
return self
|
||||||
|
|
||||||
|
sys.modules["transformers"].AutoModel = BlowsUpModel
|
||||||
|
|
||||||
with pytest.raises(RuntimeError, match="model checkpoint not found"):
|
with pytest.raises(RuntimeError, match="model checkpoint not found"):
|
||||||
local_encoder.classify_zero_shot("x", ["a"], model_id="bad-model", device="cpu")
|
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
|
||||||
Reference in New Issue
Block a user