feat(classifier): natural-language category labels + multi_label for local_encoder #50

Merged
alee merged 1 commits from feat/encoder-natural-language-labels into main 2026-09-07 02:53:58 +00:00
2 changed files with 115 additions and 3 deletions

View File

@@ -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])

View File

@@ -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."""