diff --git a/src/local_encoder.py b/src/local_encoder.py index 7847134..274836d 100644 --- a/src/local_encoder.py +++ b/src/local_encoder.py @@ -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]) diff --git a/tests/test_local_encoder.py b/tests/test_local_encoder.py index 90c30bb..105b155 100644 --- a/tests/test_local_encoder.py +++ b/tests/test_local_encoder.py @@ -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."""