feat(classifier): natural-language category labels + multi_label for local_encoder #50
@@ -30,6 +30,38 @@ _MISSING_DEPENDENCY_MESSAGE = (
|
||||
"pip install -r requirements-encoder.txt"
|
||||
)
|
||||
|
||||
# HF's zero-shot pipeline scores each candidate label against the input using
|
||||
# a hypothesis template ("This example is {}."), so the label itself has to
|
||||
# read as natural language for the entailment scoring to work well -- feeding
|
||||
# it a raw config identifier like "tool_use_agentic" or "diff_checking" asks
|
||||
# the model to judge "This example is tool_use_agentic.", which is not a
|
||||
# sentence its NLI training ever saw. Measured live 2026-09-06 against
|
||||
# bart-large-mnli: passing cfg.proficiency.categories's raw strings directly
|
||||
# (the previous behavior) put 4 of 9 test prompts on the wrong category, with
|
||||
# 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] = {
|
||||
"coding_general": "writing new code",
|
||||
"coding_refactor": (
|
||||
"refactoring or restructuring existing code without changing its behavior"
|
||||
),
|
||||
"debugging": "finding and fixing a bug in code",
|
||||
"docs_writing": "writing documentation, comments, or explanations",
|
||||
"summarization": "summarizing or condensing a long text",
|
||||
"file_summarization": "summarizing the contents of a file",
|
||||
"diff_checking": "reviewing a code diff or comparing code changes",
|
||||
"translation": "translating text from one language to another",
|
||||
"reasoning_math": "solving a math problem or a logical reasoning puzzle",
|
||||
"tool_use_agentic": (
|
||||
"a multi-step task that requires using tools, running commands, "
|
||||
"or reading and writing files"
|
||||
),
|
||||
"general_chat": "casual conversation or a general question",
|
||||
}
|
||||
|
||||
# The loaded HF pipeline, cached at module level so it survives across
|
||||
# calls within one process — reloading a model per request would be far
|
||||
# slower than the LLM call this mode replaces. Keyed by (model_id, device)
|
||||
@@ -85,7 +117,18 @@ def classify_zero_shot(
|
||||
OOM, ...) rather than swallowing it — the caller treats any exception
|
||||
here identically to a local-LLM parse failure and falls through the
|
||||
existing cascade.
|
||||
|
||||
Categories are scored under their natural-language description
|
||||
(``_CATEGORY_DESCRIPTIONS``), never the raw config identifier -- see that
|
||||
map's docstring for why. ``multi_label=True`` scores each candidate
|
||||
independently instead of normalizing them to sum to 1: when several
|
||||
categories are plausible, the default (single-label) mode forces them to
|
||||
compete for probability mass, which drags down the correct answer's score
|
||||
even when it is a clear match on its own terms.
|
||||
"""
|
||||
classifier = _load_pipeline(model_id, device)
|
||||
result = classifier(task, candidate_labels=categories)
|
||||
return result["labels"][0], float(result["scores"][0])
|
||||
descriptions = [_CATEGORY_DESCRIPTIONS.get(c, c) for c in categories]
|
||||
desc_to_category = dict(zip(descriptions, categories))
|
||||
result = classifier(task, candidate_labels=descriptions, multi_label=True)
|
||||
top_description = result["labels"][0]
|
||||
return desc_to_category[top_description], float(result["scores"][0])
|
||||
|
||||
@@ -66,8 +66,14 @@ def test_missing_dependency_raises_an_actionable_message(monkeypatch, restore_sy
|
||||
|
||||
|
||||
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": ["coding_general", "general_chat"], "scores": [0.87, 0.13]}
|
||||
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)
|
||||
|
||||
@@ -79,6 +85,69 @@ def test_classify_zero_shot_returns_top_label_and_score(restore_sys_modules):
|
||||
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."""
|
||||
|
||||
Reference in New Issue
Block a user