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"
|
"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
|
# The loaded HF pipeline, cached at module level so it survives across
|
||||||
# calls within one process — reloading a model per request would be far
|
# 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)
|
# 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
|
OOM, ...) rather than swallowing it — the caller treats any exception
|
||||||
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
|
||||||
|
(``_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)
|
classifier = _load_pipeline(model_id, device)
|
||||||
result = classifier(task, candidate_labels=categories)
|
descriptions = [_CATEGORY_DESCRIPTIONS.get(c, c) for c in categories]
|
||||||
return result["labels"][0], float(result["scores"][0])
|
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):
|
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(
|
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)
|
_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)
|
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):
|
def test_pipeline_is_built_once_and_reused(restore_sys_modules):
|
||||||
"""Loading a model per request would be far slower than the LLM call
|
"""Loading a model per request would be far slower than the LLM call
|
||||||
this mode replaces -- the pipeline must be cached, not rebuilt."""
|
this mode replaces -- the pipeline must be cached, not rebuilt."""
|
||||||
|
|||||||
Reference in New Issue
Block a user