feat(classifier): add classifier.mode local_decision (first-token logprob classifier on local Ollama) #106
20
CLAUDE.md
20
CLAUDE.md
@@ -868,7 +868,7 @@ weeks.
|
||||
### Which implementation is PRIMARY is now a config choice
|
||||
|
||||
`classifier.mode` in the live deployment is currently `local_encoder`, set in
|
||||
`config.local.yaml` with `device: cuda` on the classifier host. That is the
|
||||
`config.local.yaml` with `device: cpu` on the classifier host. That is the
|
||||
mode answering real traffic right now. The explicit caveat is that a
|
||||
CPU-vs-CUDA latency and confidence comparison on the live classifier host is
|
||||
still pending; until that measurement exists, flipping the project default to
|
||||
@@ -903,8 +903,8 @@ probability mass), brought that to 8 of 9 correct with confidence scores
|
||||
threshold and correctly fell through to the safe fallback rather than
|
||||
mis-routing.
|
||||
|
||||
`classifier.mode` (`local_llm` default, `cloud_llm`, `local_encoder`) picks
|
||||
what answers a classification request — a peer concept to the cascade above,
|
||||
`classifier.mode` (`local_llm` default, `cloud_llm`, `local_encoder`,
|
||||
`local_decision`) picks what answers a classification request — a peer concept to the cascade above,
|
||||
**not a replacement for it**. Whichever mode is primary, a failure still
|
||||
walks the exact same cascade (stale session → session history →
|
||||
`cloud_fallback` → the static guess), unmodified.
|
||||
@@ -929,6 +929,16 @@ walks the exact same cascade (stale session → session history →
|
||||
falls back to `fallback_tier` — a real limitation, not a bug. A
|
||||
below-threshold confidence is treated as a failure and cascades exactly
|
||||
like a local-LLM parse failure would.
|
||||
- **`local_decision`** asks a small generative model (configured in
|
||||
`classifier.decision`) to pick a category from a set of natural-language
|
||||
descriptions, then optionally classifies tier and — when enabled in its
|
||||
decision config — runs the same A/B confidence check on the chosen label.
|
||||
Requires a `classifier.decision:` block in config (confidence threshold,
|
||||
base URL, model name); it is an opt-in overlay mode, not the default, and
|
||||
the config load will refuse to start without it when the mode is selected.
|
||||
Only produces `task_category` and an optional `task_tier` when the
|
||||
decision block's `tier_classification.enabled` is true. Below-threshold
|
||||
confidence cascades like a parse failure would.
|
||||
- **Neither is gated by `local_compute.enabled`** (gaming mode, below) the
|
||||
way `local_llm` is: `cloud_llm` never touches local hardware, and
|
||||
`local_encoder` is small enough to run on CPU, so neither competes for the
|
||||
@@ -943,7 +953,9 @@ value (see [admin-portal](docs/admin-portal.md)).
|
||||
The default in `config.yaml` is still `local_llm`. `local_encoder` is
|
||||
intentionally an opt-in per-deployment choice rather than the repository
|
||||
default until the pending CPU-vs-CUDA comparison on the live classifier host is
|
||||
available.
|
||||
available. `local_decision` is also opt-in: it requires a `classifier.decision:`
|
||||
block with its own base URL, model, and confidence thresholds — the config load
|
||||
refuses to start without it when the mode is selected.
|
||||
|
||||
## Local dispatch model
|
||||
|
||||
|
||||
@@ -269,6 +269,7 @@ header.navbar{padding-top:2px!important;padding-bottom:2px!important}
|
||||
<option value="local_llm">local_llm</option>
|
||||
<option value="cloud_llm">cloud_llm</option>
|
||||
<option value="local_encoder">local_encoder</option>
|
||||
<option value="local_decision">local_decision</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="col-md-8" id="classifier-mode-fields"></div>
|
||||
@@ -1183,6 +1184,54 @@ function classifierModeFieldsHtml(mode, data) {
|
||||
</div>
|
||||
</div>`;
|
||||
}
|
||||
if (mode === 'local_decision') {
|
||||
const dec = (data && data.decision) || {};
|
||||
return `
|
||||
<div class="row g-2">
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" id="classifier-decision-base-url"
|
||||
data-key="decision.base_url"
|
||||
placeholder="http://localhost:11434"
|
||||
value="${escapeHtml(dec.base_url || '')}">
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" id="classifier-decision-model"
|
||||
data-key="decision.model"
|
||||
placeholder="qwen3.5:4b"
|
||||
value="${escapeHtml(dec.model || '')}">
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" type="number" id="classifier-decision-num-ctx"
|
||||
data-key="decision.num_ctx"
|
||||
placeholder="8192"
|
||||
value="${dec.num_ctx != null ? dec.num_ctx : ''}">
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" type="number" id="classifier-decision-timeout-s"
|
||||
data-key="decision.timeout_s"
|
||||
placeholder="10"
|
||||
value="${dec.timeout_s != null ? dec.timeout_s : ''}">
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" type="number" step="0.1" id="classifier-decision-confidence-min"
|
||||
data-key="decision.confidence_min"
|
||||
placeholder="0.5"
|
||||
value="${dec.confidence_min != null ? dec.confidence_min : ''}">
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<input class="form-control form-control-sm" type="number" step="0.1" id="classifier-decision-coverage-min"
|
||||
data-key="decision.coverage_min"
|
||||
placeholder="0.3"
|
||||
value="${dec.coverage_min != null ? dec.coverage_min : ''}">
|
||||
</div>
|
||||
</div>
|
||||
<div class="form-check form-switch mt-2">
|
||||
<input class="form-check-input" type="checkbox" id="classifier-decision-tier-enabled"
|
||||
data-key="decision.tier_enabled"
|
||||
${dec.tier_enabled ? 'checked' : ''}>
|
||||
<label class="form-check-label" for="classifier-decision-tier-enabled" style="font-size:.78rem">tier_enabled</label>
|
||||
</div>`;
|
||||
}
|
||||
return `<span class="text-muted" style="font-size:.78rem">Today's default: the local model configured under <code>classifier.model</code> above.</span>`;
|
||||
}
|
||||
|
||||
@@ -1285,6 +1334,16 @@ async function loadClassifierConfig() {
|
||||
const badge = document.getElementById('classifier-mode-badge');
|
||||
badge.textContent = data.mode + (data.mode_source === 'overlay' ? ' (overlay)' : '');
|
||||
badge.className = data.mode === 'local_llm' ? 'badge bg-secondary' : 'badge bg-info';
|
||||
if (data.mode === 'local_decision' && data.decision) {
|
||||
const d = data.decision;
|
||||
if (d.base_url) document.getElementById('classifier-decision-base-url').value = d.base_url;
|
||||
if (d.model) document.getElementById('classifier-decision-model').value = d.model;
|
||||
if (d.num_ctx) document.getElementById('classifier-decision-num-ctx').value = d.num_ctx;
|
||||
if (d.timeout_s) document.getElementById('classifier-decision-timeout-s').value = d.timeout_s;
|
||||
if (d.confidence_min) document.getElementById('classifier-decision-confidence-min').value = d.confidence_min;
|
||||
if (d.coverage_min) document.getElementById('classifier-decision-coverage-min').value = d.coverage_min;
|
||||
document.getElementById('classifier-decision-tier-enabled').checked = !!d.tier_enabled;
|
||||
}
|
||||
}
|
||||
|
||||
function collectClassifierConfigBody() {
|
||||
@@ -1317,6 +1376,16 @@ function collectClassifierConfigBody() {
|
||||
// wants the 0.0-1.0 probability classify_zero_shot actually returns.
|
||||
if (thresholdPct !== '') encoder.confidence_min = parseFloat(thresholdPct) / 100;
|
||||
body.encoder = encoder;
|
||||
} else if (mode === 'local_decision') {
|
||||
body.decision = {
|
||||
base_url: document.getElementById('classifier-decision-base-url')?.value || 'http://localhost:11434',
|
||||
model: document.getElementById('classifier-decision-model')?.value || 'qwen3.5:4b',
|
||||
num_ctx: parseInt(document.getElementById('classifier-decision-num-ctx')?.value) || 8192,
|
||||
timeout_s: parseInt(document.getElementById('classifier-decision-timeout-s')?.value) || 10,
|
||||
confidence_min: parseFloat(document.getElementById('classifier-decision-confidence-min')?.value) || 0.5,
|
||||
coverage_min: parseFloat(document.getElementById('classifier-decision-coverage-min')?.value) || 0.3,
|
||||
tier_enabled: document.getElementById('classifier-decision-tier-enabled')?.checked || false,
|
||||
};
|
||||
}
|
||||
return body;
|
||||
}
|
||||
|
||||
@@ -680,22 +680,22 @@ local_dispatch_models:
|
||||
# Q4_K_M @ 32k: 1.14s mean. Q8_0 @ 32k: 3.79s, because it needs 24GB and
|
||||
# Ollama spills 17% to CPU (watch the "17%/83%" column in `ollama ps`).
|
||||
#
|
||||
# num_ctx 32768 costs ~4GB of KV cache over 16384 (13GB -> 17GB resident).
|
||||
# That is affordable here ONLY because the vision model was retagged to
|
||||
# num_ctx 8192; together they are 23.2GB of 24GB, with ~800MB spare. On a
|
||||
# smaller card, drop this to 16384 first -- KV cache is the cheapest GB to
|
||||
# reclaim, and the classifier never needs it (classifier.max_input_chars
|
||||
# clamps to 8000 chars, ~2.7k tokens).
|
||||
# num_ctx 16384 measured (see plans/local-decision-classifier-results.md
|
||||
# sections 7 and 9): the 14b at 32k is 16.5-17.8 GB resident, at 16k it is
|
||||
# 12.26 GB. That headroom is what lets the local_decision classifier
|
||||
# qwen3.5:4b (num_ctx 8192, 5.9 GB) coexist: 12.3 + 5.9 = ~19.9 of 24 GB
|
||||
# (measured 2026-09-28). Local vision is off because its model no longer
|
||||
# fits.
|
||||
#
|
||||
# Modelfile recipe (run on an Ollama host):
|
||||
# ollama pull qwen2.5-coder:14b
|
||||
# printf 'FROM qwen2.5-coder:14b\nPARAMETER num_ctx 32768\n' > Modelfile.local-dispatch
|
||||
# printf 'FROM qwen2.5-coder:14b\nPARAMETER num_ctx 16384\n' > Modelfile.local-dispatch
|
||||
# ollama create qwen2.5-coder-router:14b -f Modelfile.local-dispatch
|
||||
model_id: "qwen2.5-coder-router:14b"
|
||||
base_url: "http://localhost:11434/v1"
|
||||
api_key_env:
|
||||
timeout_seconds: 180
|
||||
context_window: 32768
|
||||
context_window: 16384
|
||||
max_output_tokens: 2048
|
||||
tier: 1
|
||||
# Restricted to file_summarization and diff_checking. The router's hard
|
||||
|
||||
@@ -151,19 +151,32 @@ can't safely represent:
|
||||
|
||||
- **Local Compute** — "gaming mode". See [Gaming mode](#gaming-mode) below.
|
||||
- **Classifier** — `classifier.mode`'s config (`cloud_llm` needs a pinned
|
||||
primary or `auto`; `local_encoder` needs a model id). It shows the current
|
||||
mode and, for `cloud_primary_auto`, the **live** resolved cheapest
|
||||
candidate — computed with `routing.cheapest_classifier_candidate`, the same
|
||||
function `classify()` calls, not a static echo of the config. An admin
|
||||
panel that shows a config value instead of what dispatch actually does is a
|
||||
real class of bug — the same lesson the profiles zero-admit badge above
|
||||
already taught this project — and a "live" reading that is secretly stale
|
||||
is worse than no reading at all. `POST /admin/api/classifier-config`
|
||||
validates the same way config load does: an invalid combination (e.g.
|
||||
`cloud_llm` with neither a pinned primary nor `auto`) is rejected before
|
||||
anything reaches `config.local.yaml`, and mode plus its companion block are
|
||||
written as one atomic change so an in-between invalid state is never even
|
||||
written transiently.
|
||||
primary or `auto`; `local_encoder` needs a model id; `local_decision` needs a
|
||||
generative model id). It shows the current mode and, for `cloud_primary_auto`,
|
||||
the **live** resolved cheapest candidate — computed with
|
||||
`routing.cheapest_classifier_candidate`, the same function `classify()`
|
||||
calls, not a static echo of the config. An admin panel that shows a config
|
||||
value instead of what dispatch actually does is a real class of bug — the same
|
||||
lesson the profiles zero-admit badge above already taught this project — and a
|
||||
"live" reading that is secretly stale is worse than no reading at all.
|
||||
`POST /admin/api/classifier-config` validates the same way config load does:
|
||||
an invalid combination (e.g. `cloud_llm` with neither a pinned primary nor
|
||||
`auto`) is rejected before anything reaches `config.local.yaml`, and mode
|
||||
plus its companion block are written as one atomic change so an in-between
|
||||
invalid state is never even written transiently.
|
||||
|
||||
When `local_decision` is selected the panel reveals a **Local Decision**
|
||||
block with the fields from `LocalDecisionConfig` in `src/config.py`:
|
||||
|
||||
| field | type | default | meaning |
|
||||
|---|---|---|---|
|
||||
| `base_url` | string | `http://localhost:11434` | Ollama endpoint for the generative classifier |
|
||||
| `model` | string | `qwen3.5:4b` | The small generative model that picks category by choice |
|
||||
| `num_ctx` | int | `8192` | Context window passed to the model |
|
||||
| `timeout_s` | int | `10` | Seconds before the request times out |
|
||||
| `confidence_min` | float | `0.5` | Minimum logprob-derived confidence to accept the verdict (range [0.0, 1.0]); below this the classification cascades as a failure |
|
||||
| `coverage_min` | float | `0.3` | Minimum total probability mass on option letters; below this `parse_logprobs` raises and the classification cascades |
|
||||
| `tier_enabled` | bool | `false` | Whether the classifier may also decide `task_tier` (vs. only category) |
|
||||
- **Watchdog** — the watchdog's live state (`controls.html#watchdog-card`,
|
||||
`loadWatchdog()`): the last tick time (and a `no tick` badge before the
|
||||
first one), how many sessions the last tick saw, how many verdicts it
|
||||
|
||||
@@ -329,6 +329,72 @@ CPU install:
|
||||
pip install -r requirements-encoder.txt
|
||||
```
|
||||
|
||||
### `local_decision`: read the first-token logprobs, never parse a generation
|
||||
|
||||
`classifier.mode: local_decision` is a third answer to the same
|
||||
runaway-reasoning problem, and the most direct one. Instead of asking a small
|
||||
generative model to emit JSON and hoping it stops, it asks for a single
|
||||
lettered choice and reads the decision out of the **logprobs on the first
|
||||
output token**.
|
||||
|
||||
Mechanism, a Jev-style first-token-logprob classifier:
|
||||
|
||||
- The task text and a lettered list of candidate categories go into one prompt;
|
||||
the model is asked to answer with one letter.
|
||||
- Ollama returns `logprobs` for that first position. `local_decision.parse_logprobs`
|
||||
sums `exp(logprob)` per option letter (A through J), ignoring non-option
|
||||
tokens and Ollama's `.` separator tokens, then normalizes the accumulated mass
|
||||
into a confidence. The letter with the most mass is the predicted category.
|
||||
- There is **no JSON parse, no reasoning trace, and no multi-token generation
|
||||
to run away**. `num_predict` is 1 and `think` is false, so a model that would
|
||||
normally spend its budget thinking has nowhere to spend it. The whole failure
|
||||
class that took `qwen3.5` down on 4 of 19 calls simply does not exist here.
|
||||
- Confidence is the chosen letter's share of total option mass. `coverage_min`
|
||||
(default `0.3`) is the minimum total option mass for a call to be accepted at
|
||||
all; below it, or when no logprobs come back, the call raises and walks the
|
||||
same cascade a `local_llm` parse failure would. `confidence_min` (default
|
||||
`0.5`) treats a below-threshold verdict as a failure rather than a
|
||||
low-confidence answer, mirroring the encoder's semantics.
|
||||
- As with `local_encoder`, it produces only `task_category` unless
|
||||
`classifier.decision.tier_enabled` is set, in which case it may also choose
|
||||
`task_tier`; otherwise tier falls back to `classifier.fallback_tier`.
|
||||
|
||||
Pull and enable it:
|
||||
|
||||
```bash
|
||||
ollama pull qwen3.5:4b
|
||||
```
|
||||
|
||||
```yaml
|
||||
# config.local.yaml
|
||||
classifier:
|
||||
mode: "local_decision"
|
||||
decision: {}
|
||||
```
|
||||
|
||||
Every field in `classifier.decision` has a default, so `decision: {}` is enough
|
||||
to opt in. The one that matters for hardware is **`classifier.decision.num_ctx`**:
|
||||
it sets the context window the classifier loads at (default `8192`) and scales
|
||||
the KV cache directly, so it is the cheapest GB to reclaim when the pair below
|
||||
does not fit. Model, `base_url`, timeout, and the confidence gates all live
|
||||
under the same block.
|
||||
|
||||
VRAM, and why this is tight on a 24 GB card:
|
||||
|
||||
| model | role | num_ctx | resident |
|
||||
|---|---|---|---|
|
||||
| `qwen2.5-coder-router:14b` | classify + verify + dispatch | 16384 | 12.3 GB |
|
||||
| `qwen3.5:4b` | `local_decision` | 8192 | ~7 GB |
|
||||
| | | **total** | **~19.9 GB of 24 GB** |
|
||||
|
||||
The pair fits, with about 4 GB spare, but **only because the 14b router runs at
|
||||
`num_ctx 16384`, not the 32768 it was originally tagged with**. At 32k the
|
||||
router alone measured 16.5 to 17.8 GB resident, and the two models then fit by
|
||||
about 1 GB or get evicted outright depending on load order. Cutting the router
|
||||
to 16k is what makes the coexistence robust; that is a change to the
|
||||
`qwen2.5-coder-router:14b` Modelfile, not to this config block. Local vision is
|
||||
off in this profile because its 5.6 GB model no longer fits alongside both.
|
||||
|
||||
### The latency tax, and what has and hasn't addressed it
|
||||
|
||||
~10s of local overhead per message was the original number on a reasoning
|
||||
|
||||
31
evals/heldout.yaml
Normal file
31
evals/heldout.yaml
Normal file
@@ -0,0 +1,31 @@
|
||||
tasks:
|
||||
- {id: h1, category: coding_general, prompt: "Add a /healthz endpoint to the FastAPI app that returns the git sha and uptime in seconds."}
|
||||
- {id: h2, category: coding_general, prompt: "write me a bash script that watches ~/Downloads and moves any .pdf into ~/Documents/pdfs, creating the folder if needed"}
|
||||
- {id: h3, category: coding_general, prompt: "Implement a React hook useDebounce(value, delayMs) with a small test."}
|
||||
- {id: h4, category: coding_refactor, prompt: "dispatcher.py is 2700 lines. Split the classifier cascade out into its own module without changing behaviour; keep the public function names."}
|
||||
- {id: h5, category: coding_refactor, prompt: "Replace all these nested if/else chains in parse_config with a lookup table, same outputs."}
|
||||
- {id: h6, category: coding_refactor, prompt: "Rename the `usr` variable to `user` everywhere in the auth package and pull the duplicated token-validation code into one helper."}
|
||||
- {id: h7, category: debugging, prompt: "The service crash-loops on start with `KeyError: 'HF_HOME'` after the last deploy. Figure out why and fix it."}
|
||||
- {id: h8, category: debugging, prompt: "tests pass locally but test_router_timeout fails in CI about 1 in 5 runs. can you find the race"}
|
||||
- {id: h9, category: debugging, prompt: "Why does this return None for empty lists?\n\ndef first(xs):\n for x in xs:\n return x"}
|
||||
- {id: h10, category: docs_writing, prompt: "Write a README section explaining how to configure the classifier overlay, with an example config.local.yaml."}
|
||||
- {id: h11, category: docs_writing, prompt: "add docstrings to every public function in scoring.py, google style"}
|
||||
- {id: h12, category: docs_writing, prompt: "Draft the release notes for v0.9: new admin knobs, the circuit breaker fix, and the OpenRouter allowlist."}
|
||||
- {id: h13, category: summarization, prompt: "Here's the incident postmortem thread from Slack (pasted below). Give me a 5-bullet summary of what happened and the action items.\n\n[thread: 40 messages about a database failover at 03:12, replica lag, a bad DNS TTL, and follow-ups assigned to three people]"}
|
||||
- {id: h14, category: summarization, prompt: "TL;DR this paper abstract and intro for me in two sentences: Large language models are increasingly deployed as routers..."}
|
||||
- {id: h15, category: summarization, prompt: "Condense these meeting notes into decisions and open questions only."}
|
||||
- {id: h16, category: file_summarization, prompt: "Read src/watchdog.py and tell me what it does at a high level."}
|
||||
- {id: h17, category: file_summarization, prompt: "what's in docs/incidents.md? give me an overview of each incident"}
|
||||
- {id: h18, category: file_summarization, prompt: "Summarize the contents of config/config.yaml section by section."}
|
||||
- {id: h19, category: diff_checking, prompt: "Review this PR diff and flag anything risky:\n\n--- a/src/routing.py\n+++ b/src/routing.py\n@@ -40,7 +40,7 @@\n- if cost > ceiling:\n+ if cost >= ceiling:\n return None"}
|
||||
- {id: h20, category: diff_checking, prompt: "Compare the two versions of the migration and tell me what changed in the schema."}
|
||||
- {id: h21, category: diff_checking, prompt: "look over my staged changes before I commit, anything I forgot?"}
|
||||
- {id: h22, category: translation, prompt: "Translate the error messages in i18n/en.json into German, keep the keys."}
|
||||
- {id: h23, category: translation, prompt: "¿Puedes traducir este párrafo al inglés? El enrutador clasifica cada tarea antes de enviarla."}
|
||||
- {id: h24, category: translation, prompt: "Put this email into polite Japanese for a client."}
|
||||
- {id: h25, category: reasoning_math, prompt: "If energy is billed at $8/kWh and a request uses 2.7e-5 kWh, what does 40,000 requests cost?"}
|
||||
- {id: h26, category: reasoning_math, prompt: "Three switches, one bulb in another room, you can only check once. How do you tell which switch controls it?"}
|
||||
- {id: h27, category: reasoning_math, prompt: "Prove that the sum of the first n odd numbers is n squared."}
|
||||
- {id: h28, category: general_chat, prompt: "hey, what's a good name for a model router project?"}
|
||||
- {id: h29, category: general_chat, prompt: "Is Rust or Go better for a home lab CLI tool? Just your opinion."}
|
||||
- {id: h30, category: general_chat, prompt: "thanks, that worked! what should I look at next?"}
|
||||
31
plans/local-decision-classifier-heldout.yaml
Normal file
31
plans/local-decision-classifier-heldout.yaml
Normal file
@@ -0,0 +1,31 @@
|
||||
tasks:
|
||||
- {id: h1, category: coding_general, prompt: "Add a /healthz endpoint to the FastAPI app that returns the git sha and uptime in seconds."}
|
||||
- {id: h2, category: coding_general, prompt: "write me a bash script that watches ~/Downloads and moves any .pdf into ~/Documents/pdfs, creating the folder if needed"}
|
||||
- {id: h3, category: coding_general, prompt: "Implement a React hook useDebounce(value, delayMs) with a small test."}
|
||||
- {id: h4, category: coding_refactor, prompt: "dispatcher.py is 2700 lines. Split the classifier cascade out into its own module without changing behaviour; keep the public function names."}
|
||||
- {id: h5, category: coding_refactor, prompt: "Replace all these nested if/else chains in parse_config with a lookup table, same outputs."}
|
||||
- {id: h6, category: coding_refactor, prompt: "Rename the `usr` variable to `user` everywhere in the auth package and pull the duplicated token-validation code into one helper."}
|
||||
- {id: h7, category: debugging, prompt: "The service crash-loops on start with `KeyError: 'HF_HOME'` after the last deploy. Figure out why and fix it."}
|
||||
- {id: h8, category: debugging, prompt: "tests pass locally but test_router_timeout fails in CI about 1 in 5 runs. can you find the race"}
|
||||
- {id: h9, category: debugging, prompt: "Why does this return None for empty lists?\n\ndef first(xs):\n for x in xs:\n return x"}
|
||||
- {id: h10, category: docs_writing, prompt: "Write a README section explaining how to configure the classifier overlay, with an example config.local.yaml."}
|
||||
- {id: h11, category: docs_writing, prompt: "add docstrings to every public function in scoring.py, google style"}
|
||||
- {id: h12, category: docs_writing, prompt: "Draft the release notes for v0.9: new admin knobs, the circuit breaker fix, and the OpenRouter allowlist."}
|
||||
- {id: h13, category: summarization, prompt: "Here's the incident postmortem thread from Slack (pasted below). Give me a 5-bullet summary of what happened and the action items.\n\n[thread: 40 messages about a database failover at 03:12, replica lag, a bad DNS TTL, and follow-ups assigned to three people]"}
|
||||
- {id: h14, category: summarization, prompt: "TL;DR this paper abstract and intro for me in two sentences: Large language models are increasingly deployed as routers..."}
|
||||
- {id: h15, category: summarization, prompt: "Condense these meeting notes into decisions and open questions only."}
|
||||
- {id: h16, category: file_summarization, prompt: "Read src/watchdog.py and tell me what it does at a high level."}
|
||||
- {id: h17, category: file_summarization, prompt: "what's in docs/incidents.md? give me an overview of each incident"}
|
||||
- {id: h18, category: file_summarization, prompt: "Summarize the contents of config/config.yaml section by section."}
|
||||
- {id: h19, category: diff_checking, prompt: "Review this PR diff and flag anything risky:\n\n--- a/src/routing.py\n+++ b/src/routing.py\n@@ -40,7 +40,7 @@\n- if cost > ceiling:\n+ if cost >= ceiling:\n return None"}
|
||||
- {id: h20, category: diff_checking, prompt: "Compare the two versions of the migration and tell me what changed in the schema."}
|
||||
- {id: h21, category: diff_checking, prompt: "look over my staged changes before I commit, anything I forgot?"}
|
||||
- {id: h22, category: translation, prompt: "Translate the error messages in i18n/en.json into German, keep the keys."}
|
||||
- {id: h23, category: translation, prompt: "¿Puedes traducir este párrafo al inglés? El enrutador clasifica cada tarea antes de enviarla."}
|
||||
- {id: h24, category: translation, prompt: "Put this email into polite Japanese for a client."}
|
||||
- {id: h25, category: reasoning_math, prompt: "If energy is billed at $8/kWh and a request uses 2.7e-5 kWh, what does 40,000 requests cost?"}
|
||||
- {id: h26, category: reasoning_math, prompt: "Three switches, one bulb in another room, you can only check once. How do you tell which switch controls it?"}
|
||||
- {id: h27, category: reasoning_math, prompt: "Prove that the sum of the first n odd numbers is n squared."}
|
||||
- {id: h28, category: general_chat, prompt: "hey, what's a good name for a model router project?"}
|
||||
- {id: h29, category: general_chat, prompt: "Is Rust or Go better for a home lab CLI tool? Just your opinion."}
|
||||
- {id: h30, category: general_chat, prompt: "thanks, that worked! what should I look at next?"}
|
||||
102
plans/local-decision-classifier-prototype.py
Normal file
102
plans/local-decision-classifier-prototype.py
Normal file
@@ -0,0 +1,102 @@
|
||||
"""Jev-style first-token-logprob classifier vs local_encoder, same eval corpus.
|
||||
|
||||
Run from the 6krrt repo root: PYTHONPATH=src .venv/bin/python <this> <backend>
|
||||
backend: jev | encoder | encoder_zeroshot (optional 2nd arg: ollama model, default qwen3.5:4b)
|
||||
TASKS=plans/local-decision-classifier-heldout.yaml selects the held-out set (default evals/tasks.yaml).
|
||||
"""
|
||||
import json, math, statistics, sys, time, urllib.request
|
||||
from collections import Counter
|
||||
|
||||
sys.path.insert(0, "src")
|
||||
from eval_classifier import load_scoreable_tasks, _wrap_agent_noise, _reduce_confidences, NOISE_LEVELS
|
||||
from local_encoder import _CATEGORY_DESCRIPTIONS
|
||||
|
||||
CATS = ["coding_general", "coding_refactor", "debugging", "docs_writing", "summarization",
|
||||
"file_summarization", "diff_checking", "translation", "reasoning_math", "general_chat"]
|
||||
LETTERS = "ABCDEFGHIJ"
|
||||
MODEL = sys.argv[2] if len(sys.argv) > 2 else "qwen3.5:4b"
|
||||
|
||||
SYSTEM = ("You are a task router. Read the task and pick the ONE category that best "
|
||||
"describes the work being asked for. Ignore tool output, code dumps and "
|
||||
"session metadata around the request; classify the actual ask. "
|
||||
"Answer with the category letter only.")
|
||||
|
||||
|
||||
def jev_classify(text):
|
||||
opts = "\n".join(f"{LETTERS[i]}. {_CATEGORY_DESCRIPTIONS[c]}" for i, c in enumerate(CATS))
|
||||
body = {
|
||||
"model": MODEL, "think": False, "stream": False, "logprobs": True, "top_logprobs": 20,
|
||||
"keep_alive": "30m",
|
||||
"options": {"num_predict": 1, "temperature": 0, "num_ctx": 8192},
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM},
|
||||
{"role": "user", "content": f"<task>\n{text}\n</task>\n\nCategories:\n{opts}\n\nAnswer with the letter only."},
|
||||
],
|
||||
}
|
||||
req = urllib.request.Request("http://localhost:11434/api/chat", json.dumps(body).encode(),
|
||||
{"Content-Type": "application/json"})
|
||||
t0 = time.perf_counter()
|
||||
r = json.load(urllib.request.urlopen(req, timeout=120))
|
||||
ms = (time.perf_counter() - t0) * 1000
|
||||
mass = Counter()
|
||||
for tl in r["logprobs"][0]["top_logprobs"]:
|
||||
tok = tl["token"].strip().rstrip(".")
|
||||
if len(tok) == 1 and tok in LETTERS:
|
||||
mass[tok] += math.exp(tl["logprob"])
|
||||
total = sum(mass.values())
|
||||
if not total:
|
||||
return "general_chat", 0.0, ms, 0.0
|
||||
letter, p = mass.most_common(1)[0]
|
||||
# confidence = share among the supplied options; coverage = mass on any option at all
|
||||
return CATS[LETTERS.index(letter)], p / total, ms, total
|
||||
|
||||
|
||||
def main():
|
||||
backend = sys.argv[1]
|
||||
if backend.startswith("encoder"):
|
||||
import local_encoder
|
||||
if backend == "encoder_zeroshot":
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = True
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
def classify(text):
|
||||
t0 = time.perf_counter()
|
||||
cat, conf = local_encoder.classify_zero_shot(text, CATS, model_id="BAAI/bge-large-en-v1.5", device="cpu")
|
||||
return cat, conf, (time.perf_counter() - t0) * 1000, 1.0
|
||||
else:
|
||||
classify = jev_classify
|
||||
|
||||
tasks = load_scoreable_tasks(__import__("os").environ.get("TASKS", "evals/tasks.yaml"))
|
||||
classify("warm up the model")
|
||||
per_level = {l: [0, 0] for l in NOISE_LEVELS}
|
||||
lat, confs_ok, confs_bad, verdict_ok, misses = [], [], [], 0, Counter()
|
||||
for t in tasks:
|
||||
preds = []
|
||||
for level in NOISE_LEVELS:
|
||||
cat, conf, ms, _ = classify(_wrap_agent_noise(t["prompt"], level))
|
||||
lat.append(ms)
|
||||
ok = cat == t["category"]
|
||||
per_level[level][0] += ok
|
||||
per_level[level][1] += 1
|
||||
(confs_ok if ok else confs_bad).append(conf)
|
||||
if not ok:
|
||||
misses[(t["category"], cat)] += 1
|
||||
preds.append((cat, conf))
|
||||
verdict_ok += _reduce_confidences(preds)[0] == t["category"]
|
||||
|
||||
n = len(tasks)
|
||||
print(f"== {backend} {MODEL if backend == 'jev' else ''} :: {n} tasks x {len(NOISE_LEVELS)} noise levels")
|
||||
for l, (ok, tot) in per_level.items():
|
||||
print(f" {l:6s} {ok}/{tot} = {ok/tot:.1%}")
|
||||
allok = sum(v[0] for v in per_level.values())
|
||||
print(f" all {allok}/{n*3} = {allok/(n*3):.1%} majority-vote verdict {verdict_ok}/{n}")
|
||||
q = statistics.quantiles(lat, n=20)
|
||||
print(f" latency ms: p50 {statistics.median(lat):.0f} p95 {q[18]:.0f} max {max(lat):.0f}")
|
||||
if confs_ok:
|
||||
print(f" conf correct: mean {statistics.mean(confs_ok):.2f}", end="")
|
||||
if confs_bad:
|
||||
print(f" conf wrong: mean {statistics.mean(confs_bad):.2f}", end="")
|
||||
print()
|
||||
print(" top confusions (gold -> predicted):", misses.most_common(6))
|
||||
|
||||
|
||||
main()
|
||||
210
plans/local-decision-classifier-results.md
Normal file
210
plans/local-decision-classifier-results.md
Normal file
@@ -0,0 +1,210 @@
|
||||
# Gate G — local_decision measurement results
|
||||
|
||||
Status: done -- gate G measured; VRAM passes with the 14b at 16k and local vision off (section 9); accuracy and latency pass
|
||||
|
||||
Date: 2026-09-28
|
||||
Plan: `plans/local-decision-classifier.md`
|
||||
Hardware: Quadro RTX 6000, 24 GB (24576 MiB total)
|
||||
Models: `qwen3.5:4b` (classifier), `qwen2.5-coder-router:14b` (router, resident at 32768 ctx = 17.76 GB)
|
||||
Ollama: `http://localhost:11434`
|
||||
|
||||
## Headline
|
||||
|
||||
| Criterion | Measured | Threshold | Verdict |
|
||||
|---|---|---|---|
|
||||
| Eval-set accuracy (`evals/tasks.yaml`) | **97.8%** (tuned) | ≥ 85% | PASS |
|
||||
| Held-out accuracy (`evals/heldout.yaml`) | **90.0%** (tuned) | ≥ 90% | PASS (at threshold) |
|
||||
| p95 latency (14b resident) | **502 ms** | ≤ 800 ms | PASS |
|
||||
| 14b router resident after 100 calls | **EVICTED** | must stay resident | **FAIL** |
|
||||
|
||||
**Gate verdict: STOP.** The 14b router model is evicted whenever `qwen3.5:4b` is loaded —
|
||||
on this 24 GB card there is no `num_ctx` at which the two coexist. Per the plan, "STOP AND
|
||||
REPORT after G if ... the 14b router model was evicted. Do NOT start todos 4-17 in that case."
|
||||
Todos 4-17 are NOT started.
|
||||
|
||||
---
|
||||
|
||||
## 1. Accuracy per set
|
||||
|
||||
Measured with the eval harness (`--backend decision --decision-base-url http://localhost:11434`),
|
||||
majority-vote overall accuracy over the clean/short/long noise variants of each task.
|
||||
|
||||
### Baseline (original `_DECISION_DESCRIPTIONS`)
|
||||
|
||||
| Set | clean | short | long | overall |
|
||||
|---|---|---|---|---|
|
||||
| `evals/tasks.yaml` (46) | 0.739 | 0.783 | 0.761 | **0.761** |
|
||||
| `evals/heldout.yaml` (30) | 0.800 | 0.933 | 0.900 | **0.967** |
|
||||
|
||||
Baseline confusion (eval-set, dominant cells): `coding_general → debugging` (14) and
|
||||
`→ reasoning_math` (4); `file_summarization → debugging` (9). The `debugging` description
|
||||
was attracting both code-tracing and code-summary tasks.
|
||||
|
||||
### After description tuning
|
||||
|
||||
`_DECISION_DESCRIPTIONS` changes: widen `coding_general` to include "tracing what existing
|
||||
code returns" and "implementing a feature or algorithm"; narrow `debugging` to "fixing a bug";
|
||||
narrow `reasoning_math` to "pure math or logic word problem (no code)"; soften
|
||||
`file_summarization` to "summarizing or explaining".
|
||||
|
||||
| Set | clean | short | long | overall |
|
||||
|---|---|---|---|---|
|
||||
| `evals/tasks.yaml` (46) | 0.978 | 0.978 | 0.978 | **0.978** |
|
||||
| `evals/heldout.yaml` (30) | 0.800 | 0.900 | 0.833 | **0.900** |
|
||||
|
||||
Tuning lifted the eval-set from 76.1% → 97.8% (the confusion cells are gone: eval-set
|
||||
`coding_general` 27/27, `file_summarization` 18/18, `debugging` 18/18). It cost 6.7 points on
|
||||
held-out (96.7% → 90.0%), landing exactly on the 90% threshold. Held-out residual misses:
|
||||
`h2` coding_general→file_summarization, `h30` general_chat→file_summarization, `h9`
|
||||
debugging→general_chat — all low-confidence (0.33-0.49) edge cases the confidence cascade
|
||||
would absorb. Further tuning to chase these risks overfitting `evals/tasks.yaml` (the exact
|
||||
failure the plan warns about), so tuning stops here.
|
||||
|
||||
## 2. Calibration (tuned)
|
||||
|
||||
Mean confidence when correct vs incorrect (per-noise-level rows).
|
||||
|
||||
| Set | correct mean | incorrect mean |
|
||||
|---|---|---|
|
||||
| `evals/tasks.yaml` | 0.798 | 0.350 |
|
||||
| `evals/heldout.yaml` | 0.802 | 0.413 |
|
||||
|
||||
Calibrated in the right direction on both sets: the model is roughly twice as confident when
|
||||
right as when wrong. (Baseline was 0.833/0.438 eval-set, 0.815/0.300 heldout — the tuning
|
||||
moved both closer together but kept correct ≫ incorrect.)
|
||||
|
||||
## 3. Latency (p50 / p95)
|
||||
|
||||
Per-call latency of a single `classify_choice` over the eval-set, `qwen3.5:4b`, `num_ctx 8192`,
|
||||
with the 14b router resident (the production steady state).
|
||||
|
||||
```
|
||||
single category call: n=46 p50=363ms p95=502ms p99=598ms mean=379ms
|
||||
```
|
||||
|
||||
**p95 502 ms ≤ 800 ms — PASS.**
|
||||
|
||||
## 4. Open question (a) — noise isolation
|
||||
|
||||
Compare `classify_choice` on raw wrapped (noisy) text vs text passed through
|
||||
`local_encoder._isolate_task_text` first.
|
||||
|
||||
| Set | no-isolation (short+long) | isolated (short+long) | clean baseline |
|
||||
|---|---|---|---|
|
||||
| `evals/tasks.yaml` | 0.772 | 0.739 | 0.739 |
|
||||
| `evals/heldout.yaml` | 0.917 | 0.800 | 0.800 |
|
||||
|
||||
**Isolation HURTS** on both sets (−3.3 pt eval, −11.7 pt heldout). For `local_decision` the
|
||||
fenced-code/tool content `_isolate_task_text` strips is itself the signal (code-tracing,
|
||||
file-summary tasks). **Do NOT run `_isolate_task_text` before `classify_choice`; send raw text.**
|
||||
This overrides the prototype's earlier "noise cost about 2 points" note — for this backend
|
||||
isolation is strictly worse.
|
||||
|
||||
## 5. Open question (b) — tier call strategy
|
||||
|
||||
Measure separate-request latency vs a joint prompt reading two token positions.
|
||||
|
||||
| strategy | p50 | p95 | reliability |
|
||||
|---|---|---|---|
|
||||
| category call (10 options) | 361 ms | 595 ms | — |
|
||||
| tier call (3 options, separate) | 326 ms | 553 ms | — |
|
||||
| separate, sequential total | 689 ms | 1154 ms | 100% |
|
||||
| joint prompt, two token positions | 434 ms | 635 ms | **5/30 = 17%** |
|
||||
|
||||
The joint two-position read is faster (p50 434 vs 689 ms) but **unreliable**: only 5/30 (17%)
|
||||
prompts produced valid letter mass (>0.5) at BOTH token positions — the model does not
|
||||
reliably answer two questions with two letters in one generation.
|
||||
|
||||
**Recommendation: two separate calls fired in parallel** (category + tier). Parallel wall-clock
|
||||
= max(category, tier) ≈ p95 ~595 ms, fully reliable, and matches the plan's component-4 design
|
||||
("fired in parallel ... so it adds no wall-clock time"). The joint cleverness is rejected on
|
||||
reliability, not speed.
|
||||
|
||||
## 6. Open question (c) — num_ctx trade-off
|
||||
|
||||
100-call run at each `num_ctx`, `/api/ps` before/after, `qwen3.5:4b` footprint and 14b residency.
|
||||
|
||||
| num_ctx | p50 | p95 | qwen3.5:4b VRAM | 14b router resident? |
|
||||
|---|---|---|---|---|
|
||||
| 4096 | 361 ms | 588 ms | 6.15 GB | **NO** |
|
||||
| 8192 | 364 ms | 594 ms | 6.32 GB | **NO** |
|
||||
| 16384 | 364 ms | 591 ms | 6.65 GB | **NO** |
|
||||
|
||||
Latency is essentially flat across `num_ctx` (the classification prompts are short), and VRAM
|
||||
grows 6.15 → 6.65 GB. If coexistence were possible, `num_ctx 4096` would be the best
|
||||
trade-off. It is not possible: see below.
|
||||
|
||||
## 7. VRAM / router-model eviction (HARD CONSTRAINT — FAILS)
|
||||
|
||||
`/api/ps` before and after a 100-call run at `num_ctx 8192`:
|
||||
|
||||
```
|
||||
BEFORE: {'qwen2.5-coder-router:14b': 17.76}
|
||||
AFTER : {'qwen3.5:4b': 6.32} # 14b router GONE
|
||||
14b router resident after 100 calls: False
|
||||
```
|
||||
|
||||
Definitive coexistence check — 14b loaded at its production context (32768 = 17.76 GB), then a
|
||||
single `qwen3.5:4b` call at each `num_ctx`:
|
||||
|
||||
```
|
||||
after qwen3.5@4096 : {'qwen3.5:4b': 6.15} coexist? False
|
||||
after qwen3.5@8192 : {'qwen3.5:4b': 6.32} coexist? False
|
||||
after qwen3.5@16384: {'qwen3.5:4b': 6.65} coexist? False
|
||||
```
|
||||
|
||||
Why: the RTX 6000 has 24 GB. The 14b router occupies 17.76 GB (Q4_K_M @ 32768 ctx) plus ~1.2 GB
|
||||
of other GPU processes → ~18.4 GB used, ~5.5 GB free. `qwen3.5:4b` needs ≥ 6.15 GB at any
|
||||
`num_ctx`, so loading it always exceeds free VRAM and Ollama evicts the 14b router. 17.76 + 1.2 +
|
||||
6.15 = 25.1 GB > 24 GB — it cannot physically fit at `num_ctx 4096`, let alone 8192.
|
||||
|
||||
(The earlier plan note "the card sat at 23.4/24 GB with both resident at num_ctx 8192" was not
|
||||
reproducible — the 14b router here is 17.76 GB at its production 32768 context, and the two do
|
||||
not coexist.)
|
||||
|
||||
## 8. Gate summary and next action
|
||||
|
||||
- Eval-set 97.8% ≥ 85% — PASS
|
||||
- Held-out 90.0% ≥ 90% — PASS
|
||||
- p95 502 ms ≤ 800 ms — PASS
|
||||
- 14b router resident after 100 calls — **FAIL (evicted)**
|
||||
|
||||
**The local_decision approach fails gate G on VRAM and must NOT proceed to todos 4-17.**
|
||||
The accuracy and latency bars are met and the description tuning is complete, but the classifier
|
||||
cannot run on this 24 GB card without evicting the 14b router model that local dispatch depends
|
||||
on. Options to revisit before re-running gate G (out of scope here): a smaller classifier model
|
||||
that coexists in the remaining ~5.5 GB (e.g. a ~3-4 GB 2-3B model at low num_ctx), or a model
|
||||
that is small enough that 14b + classifier + other procs fit in 24 GB.
|
||||
|
||||
## 9. Re-measure with the 14b at 16k (reviewer, 2026-09-28, owner option A)
|
||||
|
||||
The owner chose option A: drop the router 14b to `num_ctx 16384` and re-measure.
|
||||
Local vision fallback (`qwen3-vl-router:4b`, 5.6 GB) is being turned off: it has
|
||||
never fired (0 route decisions have ever selected a `qwen3-vl` model).
|
||||
|
||||
Measured with a separate test tag (`FROM qwen2.5-coder:14b`, `PARAMETER num_ctx
|
||||
16384`), so the production tag was untouched. The tag was deleted afterwards.
|
||||
|
||||
| state | `/api/ps` size_vram | nvidia-smi used |
|
||||
|---|---|---|
|
||||
| 14b @ 16k alone | 12.26 GB | 14.2 GB |
|
||||
| + `qwen3.5:4b` @ 4096 | 12.26 + 5.73 GB | 19.7 GB |
|
||||
| + `qwen3.5:4b` @ 8192 | 12.26 + 5.89 GB | 19.9 GB (about 4 GB spare) |
|
||||
|
||||
**At 16k the two coexist, with about 4 GB of headroom.** The vision model (5.6 GB)
|
||||
would not also fit, which is why it is being turned off.
|
||||
|
||||
**Gate G's "evicted at every num_ctx" is not stable.** During the latency run,
|
||||
production traffic reloaded its own 14b at 32k, and it stayed resident alongside
|
||||
the 4b: 16.54 + 5.89 GB, 22.9 of 24 GB used, about 1 GB spare. Ollama's eviction
|
||||
follows its own memory estimate, which varies with load order and other GPU
|
||||
processes (gate G saw the same 14b as 17.76 GB). At 32k the pair fits only by
|
||||
about 1 GB and can flip back to eviction; 16k is the robust setting.
|
||||
|
||||
Latency, 100 calls of `classify_choice` (`num_ctx 8192`) with both models resident:
|
||||
p50 321 ms, p95 474 ms, max 2545 ms (the first cold call). The 14b stayed resident
|
||||
throughout. **p95 474 ms, PASS.**
|
||||
|
||||
**Revised gate verdict:** PASS on every criterion **provided** the 14b router tag
|
||||
runs at `num_ctx 16384` and local vision is off. Both are production changes and
|
||||
the owner's call; todos 4-17 may proceed once they are made.
|
||||
157
plans/local-decision-classifier.md
Normal file
157
plans/local-decision-classifier.md
Normal file
@@ -0,0 +1,157 @@
|
||||
# `classifier.mode: local_decision` — a Jev-style classifier on the local GPU
|
||||
|
||||
Status: planned -- plan reviewed 2026-09-28; gate G (tuning + measurement) decides whether the wiring proceeds
|
||||
Date: 2026-09-27
|
||||
Prototype: `plans/local-decision-classifier-prototype.py` (a benchmark you can run, NOT the
|
||||
implementation)
|
||||
Held-out prompts: `plans/local-decision-classifier-heldout.yaml`
|
||||
Background: TypeSafe's hosted "Jev" decision model and the OpenJev/SemIf open
|
||||
copies. We are copying the **technique**, not using either product.
|
||||
|
||||
## Goal
|
||||
|
||||
Add a fourth `classifier.mode`, `local_decision`. It asks a small local LLM
|
||||
(`qwen3.5:4b` on Ollama) one multiple-choice question and reads the answer
|
||||
from the **first output token's logprobs**, normalised over the option letters
|
||||
that were offered. There is no text generation, JSON or reasoning trace, so the
|
||||
runaway-reasoning failure that retired `local_llm` can't happen, and it is
|
||||
calibrated far better than `local_encoder`.
|
||||
|
||||
## Evidence (2026-09-27, RTX 6000 24 GB, measured with the prototype)
|
||||
|
||||
46 eval tasks and 30 held-out tasks, each at 3 noise levels (`_wrap_agent_noise`):
|
||||
|
||||
| | `evals/tasks.yaml` | held-out | p50 / p95 ms |
|
||||
|---|---|---|---|
|
||||
| `local_decision` (`qwen3.5:4b`, GPU) | 71% | 100% | 340-385 / 510-610 |
|
||||
| `local_encoder` + trained head (live, CPU) | 100%\* | 47% | 80-115 / 100-360 |
|
||||
| `local_encoder` zero-shot (CPU) | 39% | 60% | 80-115 / 100-350 |
|
||||
|
||||
\* The head was fitted on a synthetic corpus derived from `evals/tasks.yaml`, so
|
||||
this is leakage. On held-out prompts the head is **worse** than zero-shot.
|
||||
|
||||
- **Calibration:** `local_decision` averages 0.96 confidence when right and 0.71 when wrong. The encoder
|
||||
zero-shot sits at 0.2-0.3 on both.
|
||||
- **Eval-set misses:** 18 were `coding_general`→`reasoning_math` (algorithm tasks), 18 were
|
||||
`file_summarization`→`debugging`/`reasoning_math`/`docs_writing`. Both look like
|
||||
option-description problems (component 2).
|
||||
- **VRAM:** `qwen3.5:4b` holds about 7 GB at `num_ctx` 8192 next to `qwen2.5-coder-router:14b`.
|
||||
The card sat at 23.4/24 GB. It's tight, so measure again when planning.
|
||||
- **Weak held-out set:** a single author wrote all 30 prompts and they are short and clean. Treat
|
||||
100% as a ceiling, not a forecast.
|
||||
|
||||
## Mechanism (proven by the prototype)
|
||||
|
||||
Ollama 0.22 native `POST /api/chat`:
|
||||
`think: false`, `logprobs: true`, `top_logprobs: 20`,
|
||||
`options: {num_predict: 1, temperature: 0, num_ctx: N}`. The system prompt says
|
||||
"answer with the letter only". The options are `A.`–`J.` followed by a natural-language description.
|
||||
Parse `logprobs[0].top_logprobs`:
|
||||
- strip each token and remove a trailing `.`
|
||||
- sum the probability mass for each option letter (`"B"` and `" B"` both count)
|
||||
- confidence = winner's mass ÷ total option mass
|
||||
- coverage = total option mass
|
||||
|
||||
Tokens like `<think>` show up in the top 20 and must be ignored.
|
||||
|
||||
## Components
|
||||
|
||||
### 1. `src/local_decision.py` (stdlib + `requests`, no new dependency)
|
||||
|
||||
- **`classify_choice(text, options, *, base_url, model, num_ctx, timeout)`** returns
|
||||
`(label, confidence, coverage)`. It is generic over the option list, so tier
|
||||
reuses it.
|
||||
- **Coverage floor:** below a configurable `coverage_min` (the model didn't answer
|
||||
with a letter), raise. Don't guess.
|
||||
- **Pure parse function:** split the logprob parsing into its own function so tests
|
||||
can drive it with recorded responses and no Ollama.
|
||||
- **Noise:** decide whether to run `local_encoder._isolate_task_text` first. The prototype
|
||||
sent the raw noisy text and noise cost about 2 points; measure both.
|
||||
|
||||
### 2. Its own option descriptions
|
||||
|
||||
Give it a separate `_DECISION_DESCRIPTIONS` map. **Do not edit**
|
||||
`local_encoder._CATEGORY_DESCRIPTIONS`, because that would shift the encoder and
|
||||
invalidate its head. Tune it against the two confusions above, and re-measure after
|
||||
every change on **both** sets so nobody overfits `evals/tasks.yaml` the way the head did.
|
||||
|
||||
### 3. Wiring (follow how `local_encoder` was added)
|
||||
|
||||
- **`src/config.py`:** add `"local_decision"` to the `mode` Literal (line ~1187). Add a
|
||||
`classifier.decision` block (`base_url`, `model`, `num_ctx`, `timeout_s`,
|
||||
`confidence_min`, `coverage_min`, `tier_enabled`) and a validator that requires
|
||||
it when the mode is selected. Metering (line ~1712): this mode is an HTTP call to
|
||||
Ollama, so it meters like `local_llm` via its own `base_url`, not like the in-process encoder.
|
||||
- **`src/dispatcher.py`:**
|
||||
- add `_classify_via_local_decision` and a branch in `_classify_via_configured_mode`
|
||||
- add a startup readiness check in `_ensure_classifier_mode_ready` (model is pulled and
|
||||
logprobs come back)
|
||||
- on success record `source="classifier"`
|
||||
- treat below-threshold confidence as a failure that walks the **unchanged** cascade
|
||||
- **Gating:** it uses the GPU, so it obeys `_local_classifier_skip_reason()`
|
||||
(`local_compute.enabled` plus the circuit breaker), the same as `local_llm`.
|
||||
- **Admin:** every new scalar knob gets a control on the Classifier card (North Star #1).
|
||||
`tests/test_admin_knob_coverage.py` enforces this.
|
||||
|
||||
### 4. Tier (behind `tier_enabled`, default false)
|
||||
|
||||
- **How:** a second `classify_choice` call with three options (cheap/simple, general,
|
||||
frontier/high-stakes), fired **in parallel** with the category call so it adds no
|
||||
wall-clock time.
|
||||
- **Why it defaults off:** no gold tier labels exist yet. Until they do, tier keeps
|
||||
falling back to `fallback_tier`, exactly as `local_encoder` does now.
|
||||
- **Output:** a gold tier set of about 30 labelled prompts, used to decide the default.
|
||||
|
||||
### 5. Evaluation harness
|
||||
|
||||
- **Backend flag:** add a `--backend encoder|decision` flag to `src/eval_classifier.py`
|
||||
(latency is reported per call).
|
||||
- **Held-out set:** promote the held-out prompts to `evals/` and add a `--tasks` default for them.
|
||||
- **Head flag:** add a flag that runs the encoder with the head disabled, so the leakage
|
||||
stays visible.
|
||||
- **Report:** accuracy for both sets, calibration (confidence correct vs wrong), and p50/p95.
|
||||
|
||||
### 6. Docs
|
||||
|
||||
- **CLAUDE.md:** the "Which implementation is PRIMARY" section, which also still says
|
||||
the live encoder is `device: cuda` while the overlay says `cpu`. Fix that.
|
||||
- **Also:** `docs/local-models.md` and `docs/admin-portal.md`.
|
||||
|
||||
## Acceptance
|
||||
|
||||
- **Held-out accuracy:** at least 90% on the held-out set.
|
||||
- **Eval-set accuracy:** at least 85% on `evals/tasks.yaml` after component 2. **Stop and
|
||||
report** if that needs more than description tuning.
|
||||
- **Latency:** p95 at or below 800 ms on the RTX 6000 with the 14b router model resident.
|
||||
The router model must not be evicted: check `/api/ps` before and after a 100-call run.
|
||||
- **Failure paths:** Ollama down, model missing, coverage below the floor and low
|
||||
confidence each walk the cascade, covered by tests with a fake Ollama.
|
||||
- **Tests:** the full suite is green, including `test_admin_knob_coverage.py`.
|
||||
|
||||
## Out of scope
|
||||
|
||||
- **Replacing gaming mode** with an automatic Ollama reachability check. It's a separate
|
||||
lift, and when it lands this mode inherits it for free.
|
||||
- **What to do about the trained encoder head.** It's a separate decision, but the
|
||||
evidence above is a strong argument for turning it off.
|
||||
- **Other routes:** hosted Jev (cloud-only, bans benchmarking, sends raw task text off-box),
|
||||
OpenJev 27B (no FP8 on Turing, CC BY-NC, evicts the router model), and fine-tuning.
|
||||
- **Defaults:** making `local_decision` the `config.yaml` default. It's an overlay opt-in
|
||||
until it has run on live traffic.
|
||||
|
||||
## Guardrails
|
||||
|
||||
- **Worktree:** do the work in its own worktree (North Star #3). There are 15 open ones,
|
||||
and `feat/local-encoder-accuracy-rebuild` touches classifier code. Commit by explicit path.
|
||||
- **Live router:** exercise the admin portal on the 8081 sandbox, never on production 8080,
|
||||
and don't edit `config.local.yaml` (see the `6krrt-ops` skill).
|
||||
- **Live database:** open `router.db` read-only.
|
||||
|
||||
## Open questions for planning
|
||||
|
||||
1. Isolate noise before calling or not (component 1)? Decide from a measurement.
|
||||
2. Should tier's second call be a separate request, or both questions in one prompt
|
||||
read from two token positions? Separate is simpler; measure the latency before
|
||||
choosing the clever option.
|
||||
3. What `num_ctx` gives the best VRAM trade-off against the 8000-char
|
||||
`max_input_chars` clamp?
|
||||
@@ -513,6 +513,7 @@ _CONFIG_ALLOWLIST: dict[str, tuple[str, ...]] = {
|
||||
"session_cache.staleness_seconds": ("session_cache", "staleness_seconds"),
|
||||
"local_compute.enabled": ("local_compute", "enabled"),
|
||||
"verification.local_llm_enabled": ("verification", "local_llm_enabled"),
|
||||
"local_vision.enabled": ("local_vision", "enabled"),
|
||||
"pinch.enabled": ("pinch", "enabled"),
|
||||
"pinch.prefix_probe": ("pinch", "prefix_probe"),
|
||||
"pinch.relevance.enabled": ("pinch", "relevance", "enabled"),
|
||||
@@ -543,6 +544,7 @@ _CONFIG_GET_ORDER: list[str] = [
|
||||
"session_cache.staleness_seconds",
|
||||
"local_compute.enabled",
|
||||
"verification.local_llm_enabled",
|
||||
"local_vision.enabled",
|
||||
"pinch.enabled",
|
||||
"pinch.prefix_probe",
|
||||
"pinch.relevance.enabled",
|
||||
@@ -904,10 +906,11 @@ class _ClassifierConfigBody(_AdminBody):
|
||||
classes' field definitions here.
|
||||
"""
|
||||
|
||||
mode: Literal["local_llm", "cloud_llm", "local_encoder"]
|
||||
mode: Literal["local_llm", "cloud_llm", "local_encoder", "local_decision"]
|
||||
cloud_primary: Optional[dict] = None
|
||||
cloud_primary_auto: bool = False
|
||||
encoder: Optional[dict] = None
|
||||
decision: Optional[dict] = None
|
||||
|
||||
|
||||
class _CloudFallbackBody(_AdminBody):
|
||||
@@ -2671,6 +2674,7 @@ def build_router(
|
||||
"cloud_primary_auto": bool(classifier.get("cloud_primary_auto", False)),
|
||||
"cloud_fallback": classifier.get("cloud_fallback"),
|
||||
"encoder": classifier.get("encoder"),
|
||||
"decision": classifier.get("decision"),
|
||||
"resolved_primary": resolved_primary,
|
||||
}
|
||||
|
||||
@@ -2737,6 +2741,8 @@ def build_router(
|
||||
mutations.append((("classifier", "cloud_primary"), body.cloud_primary))
|
||||
if body.mode == "local_encoder":
|
||||
mutations.append((("classifier", "encoder"), body.encoder or {}))
|
||||
if body.mode == "local_decision":
|
||||
mutations.append((("classifier", "decision"), body.decision or {}))
|
||||
|
||||
try:
|
||||
_persist_many_to(
|
||||
|
||||
@@ -1119,6 +1119,42 @@ class LocalEncoderConfig(StrictModel):
|
||||
tier_feature_fields: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class LocalDecisionConfig(StrictModel):
|
||||
"""A generative local model that picks a task category by choice.
|
||||
|
||||
Unlike the embedding+centroid ``LocalEncoderConfig``, this is a small
|
||||
generative LLM asked to return one of a fixed set of category labels
|
||||
(see ``classify_choice`` in local_decision.py). Confidence comes from the
|
||||
model's logprobs on the chosen label, not from a similarity score.
|
||||
"""
|
||||
|
||||
base_url: str = "http://localhost:11434"
|
||||
model: str = "qwen3.5:4b"
|
||||
num_ctx: int = 8192
|
||||
timeout_s: int = 10
|
||||
# minimum logprob-derived confidence to accept the model's verdict.
|
||||
# Below this threshold the classification is treated as a FAILURE
|
||||
# rather than a low-confidence answer, mirroring the encoder's
|
||||
# confidence_min semantics.
|
||||
confidence_min: float = 0.5
|
||||
# minimum total probability mass on option letters for one call;
|
||||
# below this threshold parse_logprobs raises RuntimeError and the
|
||||
# classification cascades as a failure.
|
||||
coverage_min: float = 0.3
|
||||
# whether the classifier may also decide task_tier (vs. only category).
|
||||
tier_enabled: bool = False
|
||||
|
||||
@field_validator("confidence_min")
|
||||
@classmethod
|
||||
def confidence_min_in_range(cls, v: float) -> float:
|
||||
if not (0.0 <= v <= 1.0):
|
||||
raise ValueError(
|
||||
f"classifier.decision.confidence_min must be in [0.0, 1.0], "
|
||||
f"got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
class ClassifierConfig(StrictModel):
|
||||
provider: str
|
||||
base_url: str
|
||||
@@ -1184,7 +1220,7 @@ class ClassifierConfig(StrictModel):
|
||||
#
|
||||
# "local_llm" is every existing deployment's behavior, unchanged, and
|
||||
# stays the default so this feature costs nothing until opted into.
|
||||
mode: Literal["local_llm", "cloud_llm", "local_encoder"] = "local_llm"
|
||||
mode: Literal["local_llm", "cloud_llm", "local_encoder", "local_decision"] = "local_llm"
|
||||
# Only read when mode == "cloud_llm", and mutually exclusive with
|
||||
# cloud_primary_auto (RouterConfig validator enforces exactly one).
|
||||
# Reuses CloudFallbackConfig's shape verbatim rather than a near-copy —
|
||||
@@ -1197,6 +1233,8 @@ class ClassifierConfig(StrictModel):
|
||||
cloud_primary_auto: bool = False
|
||||
# Only read when mode == "local_encoder".
|
||||
encoder: Optional["LocalEncoderConfig"] = None
|
||||
# Only read when mode == "local_decision".
|
||||
decision: Optional["LocalDecisionConfig"] = None
|
||||
|
||||
|
||||
class LocalComputeConfig(StrictModel):
|
||||
@@ -1518,6 +1556,22 @@ class RouterConfig(StrictModel):
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def local_decision_mode_needs_decision_config(self) -> "RouterConfig":
|
||||
"""classifier.mode == 'local_decision' needs to know which model.
|
||||
|
||||
LocalDecisionConfig ships sensible defaults, so this never forces an
|
||||
operator to write out every field -- but the block itself must exist,
|
||||
the same way local_encoder.py refuses to guess a model id.
|
||||
"""
|
||||
if self.classifier.mode == "local_decision" and self.classifier.decision is None:
|
||||
raise ValueError(
|
||||
"classifier.mode is 'local_decision' but classifier.decision is not "
|
||||
"set. Add a classifier.decision block (its fields all have "
|
||||
"defaults, so `classifier.decision: {}` is enough)."
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def gaming_mode_requires_a_cloud_classifier(self) -> "RouterConfig":
|
||||
"""Local compute off with no cloud classifier is a worse router.
|
||||
@@ -1724,18 +1778,26 @@ class RouterConfig(StrictModel):
|
||||
# silently stopped metering a GPU whose electricity is on this
|
||||
# machine's own bill -- while warning about a URL nothing calls.
|
||||
encoder_is_in_process = self.classifier.mode == "local_encoder"
|
||||
sites = {
|
||||
"classify": (
|
||||
True if encoder_is_in_process
|
||||
classify_is_loopback = True if encoder_is_in_process else (
|
||||
_is_loopback_host(self.classifier.decision.base_url)
|
||||
if self.classifier.mode == "local_decision"
|
||||
and self.classifier.decision is not None
|
||||
else _is_loopback_host(self.classifier.base_url)
|
||||
),
|
||||
)
|
||||
sites = {
|
||||
"classify": classify_is_loopback,
|
||||
"verify": _is_loopback_host(self.verification.base_url),
|
||||
"local_vision": _is_loopback_host(self.local_vision.base_url),
|
||||
}
|
||||
for name, loopback in sites.items():
|
||||
if not loopback:
|
||||
section = {
|
||||
"classify": self.classifier,
|
||||
"classify": (
|
||||
self.classifier.decision
|
||||
if self.classifier.mode == "local_decision"
|
||||
and self.classifier.decision is not None
|
||||
else self.classifier
|
||||
),
|
||||
"verify": self.verification,
|
||||
"local_vision": self.local_vision,
|
||||
}[name]
|
||||
|
||||
@@ -61,6 +61,7 @@ import admin
|
||||
import circuit_breaker
|
||||
import events
|
||||
import exploration
|
||||
import local_decision
|
||||
import local_encoder
|
||||
import local_energy
|
||||
import logs
|
||||
@@ -570,14 +571,39 @@ def _ensure_classifier_mode_ready() -> None:
|
||||
not at the worst possible moment" reasoning StrictModel and this
|
||||
project's config validators already apply.
|
||||
"""
|
||||
if cfg.classifier.mode != "local_encoder":
|
||||
return
|
||||
if cfg.classifier.mode == "local_encoder":
|
||||
try:
|
||||
local_encoder.ensure_available(cfg.classifier.encoder.model, cfg.classifier.encoder.device)
|
||||
except ImportError as exc:
|
||||
logs.warning("classifier_mode_unavailable", mode="local_encoder", error=str(exc))
|
||||
raise
|
||||
|
||||
if cfg.classifier.mode == "local_decision":
|
||||
if not cfg.local_compute.enabled:
|
||||
logs.info("classifier_mode_probe_skipped", mode="local_decision", reason="gaming_mode")
|
||||
return
|
||||
|
||||
try:
|
||||
# Make a single test call to verify the model responds with logprobs.
|
||||
# We don't raise here because Ollama-down at boot time should not
|
||||
# stop the router from starting — the cascade handles classification
|
||||
# failures gracefully.
|
||||
local_decision.classify_category(
|
||||
"test task",
|
||||
base_url=cfg.classifier.decision.base_url,
|
||||
model=cfg.classifier.decision.model,
|
||||
num_ctx=cfg.classifier.decision.num_ctx,
|
||||
timeout_s=cfg.classifier.decision.timeout_s,
|
||||
coverage_min=cfg.classifier.decision.coverage_min,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — Ollama may be down at boot
|
||||
logs.warning(
|
||||
"classifier_mode_unavailable",
|
||||
mode="local_decision",
|
||||
error=str(exc),
|
||||
suggestion="start Ollama or check classifier.decision.base_url",
|
||||
)
|
||||
|
||||
|
||||
_ensure_classifier_mode_ready()
|
||||
|
||||
@@ -1058,6 +1084,216 @@ def _classify_via_local_encoder(task: str) -> Classification:
|
||||
)
|
||||
|
||||
|
||||
def _classify_via_local_decision(
|
||||
system_prompt: str, user_content: str
|
||||
) -> Classification:
|
||||
"""classifier.mode == "local_decision": a generative model picks a category.
|
||||
|
||||
Unlike local_encoder (which scores the raw task text against every
|
||||
category zero-shot), local_decision asks a generative Ollama model to
|
||||
choose ONE lettered category, reading the confidence off the first-token
|
||||
logprobs. When ``cfg.classifier.decision.tier_enabled`` is True the
|
||||
function also fires a second call to classify task_tier; otherwise
|
||||
tier falls back to ``cfg.classifier.fallback_tier`` with no extra call.
|
||||
required_context_tokens is 0 for the same reason as the encoder:
|
||||
chat_completions already measures the real conversation with
|
||||
estimate_prompt_tokens and takes the larger value.
|
||||
|
||||
``system_prompt`` is accepted for signature symmetry with the other
|
||||
``_classify_via_*`` functions, but local_decision supplies its own system
|
||||
prompt (see local_decision.py); only ``user_content`` (the framed task
|
||||
text) is classified. A below-threshold confidence is treated as a FAILURE
|
||||
-- raising here means the caller cascades, exactly as it would for a
|
||||
local-LLM parse failure or a low-confidence encoder verdict.
|
||||
"""
|
||||
skip = _local_classifier_skip_reason()
|
||||
if skip is not None:
|
||||
logs.warning("classify_local_skipped", reason=skip)
|
||||
raise _ClassifierSkipped(skip)
|
||||
|
||||
dec = cfg.classifier.decision
|
||||
if cfg.local_energy.enabled and cfg.local_energy_call_sites.get("classify"):
|
||||
try:
|
||||
with local_energy.measure(
|
||||
sample_interval_seconds=cfg.local_energy.sample_interval_seconds,
|
||||
sampler=local_energy.sample_nvidia_smi,
|
||||
) as measurement:
|
||||
tier = cfg.classifier.fallback_tier
|
||||
if dec.tier_enabled:
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
cat_future = pool.submit(
|
||||
local_decision.classify_category,
|
||||
user_content,
|
||||
base_url=dec.base_url,
|
||||
model=dec.model,
|
||||
num_ctx=dec.num_ctx,
|
||||
timeout_s=dec.timeout_s,
|
||||
coverage_min=dec.coverage_min,
|
||||
)
|
||||
tier_future = pool.submit(
|
||||
_classify_tier_via_local_decision, user_content
|
||||
)
|
||||
category, confidence, _coverage = cat_future.result()
|
||||
if confidence < dec.confidence_min:
|
||||
raise RuntimeError(
|
||||
f"local_decision confidence {confidence:.3f} "
|
||||
f"is below classifier.decision.confidence_min "
|
||||
f"({dec.confidence_min})"
|
||||
)
|
||||
try:
|
||||
tier = tier_future.result()
|
||||
except Exception:
|
||||
tier = cfg.classifier.fallback_tier
|
||||
else:
|
||||
category, confidence, _coverage = (
|
||||
local_decision.classify_category(
|
||||
user_content,
|
||||
base_url=dec.base_url,
|
||||
model=dec.model,
|
||||
num_ctx=dec.num_ctx,
|
||||
timeout_s=dec.timeout_s,
|
||||
coverage_min=dec.coverage_min,
|
||||
)
|
||||
)
|
||||
if confidence < dec.confidence_min:
|
||||
raise RuntimeError(
|
||||
f"local_decision confidence {confidence:.3f} "
|
||||
f"is below classifier.decision.confidence_min "
|
||||
f"({dec.confidence_min})"
|
||||
)
|
||||
|
||||
_log_local_energy(
|
||||
model_id=dec.model,
|
||||
call_type="classify",
|
||||
measurement=measurement,
|
||||
)
|
||||
return Classification(
|
||||
task_category=category,
|
||||
task_tier=tier,
|
||||
required_context_tokens=0,
|
||||
confidence=confidence,
|
||||
source="classifier",
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — meter the real draw on ANY failure
|
||||
if isinstance(exc, requests.RequestException):
|
||||
_record_failure()
|
||||
_log_local_energy(
|
||||
model_id=dec.model,
|
||||
call_type="classify",
|
||||
measurement=measurement,
|
||||
)
|
||||
raise
|
||||
# Run category and tier classification in parallel when tier_enabled.
|
||||
# NOTE: RuntimeError (below confidence_min) and other non-transport errors
|
||||
# must NOT call _record_failure() — the endpoint answered; opening the
|
||||
# circuit on every unsure classification would skip the classifier for
|
||||
# cooldown_seconds on healthy traffic.
|
||||
try:
|
||||
tier = cfg.classifier.fallback_tier
|
||||
if dec.tier_enabled:
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
cat_future = pool.submit(
|
||||
local_decision.classify_category,
|
||||
user_content,
|
||||
base_url=dec.base_url,
|
||||
model=dec.model,
|
||||
num_ctx=dec.num_ctx,
|
||||
timeout_s=dec.timeout_s,
|
||||
coverage_min=dec.coverage_min,
|
||||
)
|
||||
tier_future = pool.submit(
|
||||
_classify_tier_via_local_decision, user_content
|
||||
)
|
||||
category, confidence, _coverage = cat_future.result()
|
||||
if confidence < dec.confidence_min:
|
||||
raise RuntimeError(
|
||||
f"local_decision confidence {confidence:.3f} is below "
|
||||
f"classifier.decision.confidence_min ({dec.confidence_min})"
|
||||
)
|
||||
try:
|
||||
tier = tier_future.result()
|
||||
except Exception:
|
||||
tier = cfg.classifier.fallback_tier
|
||||
else:
|
||||
category, confidence, _coverage = local_decision.classify_category(
|
||||
user_content,
|
||||
base_url=dec.base_url,
|
||||
model=dec.model,
|
||||
num_ctx=dec.num_ctx,
|
||||
timeout_s=dec.timeout_s,
|
||||
coverage_min=dec.coverage_min,
|
||||
)
|
||||
if confidence < dec.confidence_min:
|
||||
raise RuntimeError(
|
||||
f"local_decision confidence {confidence:.3f} is below "
|
||||
f"classifier.decision.confidence_min ({dec.confidence_min})"
|
||||
)
|
||||
tier = _classify_tier_via_local_decision(user_content)
|
||||
return Classification(
|
||||
task_category=category,
|
||||
task_tier=tier,
|
||||
required_context_tokens=0,
|
||||
confidence=confidence,
|
||||
source="classifier",
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
if isinstance(exc, requests.RequestException):
|
||||
_record_failure()
|
||||
raise
|
||||
|
||||
|
||||
# Tier classification options for local_decision mode.
|
||||
# The option letters map to tiers: A→1, B→2, C→3.
|
||||
_TIER_OPTIONS: Final[dict[str, str]] = {
|
||||
"A": "simple or small task (tier 1)",
|
||||
"B": "medium complexity task (tier 2)",
|
||||
"C": "complex or hard task (tier 3)",
|
||||
}
|
||||
|
||||
|
||||
def _classify_tier_via_local_decision(task_text: str) -> int:
|
||||
"""Ask the local decision model to classify task tier via Ollama.
|
||||
|
||||
When ``cfg.classifier.decision.tier_enabled`` is False the caller
|
||||
skips tier classification entirely and returns
|
||||
``cfg.classifier.fallback_tier`` (fast path). When enabled the
|
||||
function calls ``local_decision.classify_choice()`` with the A/B/C
|
||||
tier prompt; on any network or parse failure it falls back to
|
||||
``cfg.classifier.fallback_tier``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
1, 2, or 3 for the classified tier, or ``fallback_tier`` on
|
||||
failure.
|
||||
"""
|
||||
if not cfg.classifier.decision.tier_enabled:
|
||||
return cfg.classifier.fallback_tier
|
||||
|
||||
dec = cfg.classifier.decision
|
||||
try:
|
||||
tier_label, _confidence, _coverage = local_decision.classify_choice(
|
||||
task_text,
|
||||
_TIER_OPTIONS,
|
||||
base_url=dec.base_url,
|
||||
model=dec.model,
|
||||
num_ctx=dec.num_ctx,
|
||||
timeout_s=dec.timeout_s,
|
||||
coverage_min=dec.coverage_min,
|
||||
)
|
||||
# Map option letter to tier: A→1, B→2, C→3
|
||||
letter = tier_label.strip().upper()
|
||||
mapping = {"A": 1, "B": 2, "C": 3}
|
||||
return mapping.get(letter, cfg.classifier.fallback_tier)
|
||||
except Exception: # noqa: BLE001
|
||||
logs.warning("classify_tier_failed", fallback=cfg.classifier.fallback_tier)
|
||||
return cfg.classifier.fallback_tier
|
||||
|
||||
|
||||
def _classify_via_configured_mode(
|
||||
task: str, system_prompt: str, user_content: str
|
||||
) -> Classification:
|
||||
@@ -1073,6 +1309,8 @@ def _classify_via_configured_mode(
|
||||
return _classify_via_cloud_llm(system_prompt, user_content)
|
||||
if mode == "local_encoder":
|
||||
return _classify_via_local_encoder(task)
|
||||
if mode == "local_decision":
|
||||
return _classify_via_local_decision(system_prompt, user_content)
|
||||
return _classify_via_local_llm(system_prompt, user_content)
|
||||
|
||||
|
||||
@@ -1391,6 +1629,7 @@ def classify(task: str, context: Optional[str]) -> Classification:
|
||||
TypeError,
|
||||
json.JSONDecodeError,
|
||||
RuntimeError,
|
||||
requests.RequestException,
|
||||
) as e:
|
||||
# A primary classifier that is slow, restarting, unreachable, or (for
|
||||
# local_encoder) unsure of itself must not take the caller down with
|
||||
|
||||
@@ -36,6 +36,8 @@ plan-only run needs neither package present.
|
||||
python eval_classifier.py # every task, 3 noise levels
|
||||
python eval_classifier.py --noise clean # clean prompts only
|
||||
python eval_classifier.py --categories coding_general,diff_checking
|
||||
python eval_classifier.py --backend decision # local_decision backend
|
||||
python eval_classifier.py --head-off # encoder, trained head disabled
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -52,7 +54,7 @@ import yaml
|
||||
|
||||
from config import load_config
|
||||
|
||||
TASKS_PATH = "evals/tasks.yaml"
|
||||
TASKS_PATH = "evals/heldout.yaml"
|
||||
|
||||
# The eval measures the scoreable task set. Tool-use tasks are excluded: they
|
||||
# are scored structurally by eval_proficiency and their category is
|
||||
@@ -218,6 +220,22 @@ def _reduce_confidences(preds: list[tuple[str, float]]) -> tuple[str, float]:
|
||||
return best_cat, _mean(votes[best_cat]) or 0.0
|
||||
|
||||
|
||||
def _disable_trainable_head() -> None:
|
||||
"""Force the encoder onto its zero-shot path for the rest of the process.
|
||||
|
||||
``local_encoder._get_trainable_head`` memoises the fitted logistic-regression
|
||||
head once per process. Setting ``_TRAINABLE_HEAD_RESOLVED = True`` with
|
||||
``_TRAINABLE_HEAD = None`` makes the next call skip the filesystem load and
|
||||
resolve to no head, so the encoder classifies against the plain centroid
|
||||
similarities. This is the ``--head-off`` escape hatch that keeps the
|
||||
(leaky) trained-head score visible against zero-shot on the same corpus.
|
||||
"""
|
||||
import local_encoder
|
||||
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = True
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
|
||||
|
||||
def run_eval(
|
||||
tasks: list[dict],
|
||||
categories: list[str],
|
||||
@@ -225,10 +243,31 @@ def run_eval(
|
||||
model_id: str,
|
||||
device: str,
|
||||
noise_levels: tuple[str, ...],
|
||||
backend: str = "encoder",
|
||||
head_off: bool = False,
|
||||
decision: Optional[dict] = None,
|
||||
) -> dict:
|
||||
"""Run the encoder over *tasks* and return the full result structure."""
|
||||
"""Run the chosen classifier backend over *tasks*.
|
||||
|
||||
``backend == "encoder"`` (default) drives ``local_encoder.classify_zero_shot``;
|
||||
when ``head_off`` is set the trained head is disabled first (see
|
||||
``_disable_trainable_head``). ``backend == "decision"`` drives
|
||||
``local_decision.classify_choice`` over ``_DECISION_DESCRIPTIONS``, reading
|
||||
its ``base_url``/``model``/``num_ctx``/``timeout_s``/``coverage_min`` from
|
||||
the ``decision`` dict. ``local_decision`` is imported lazily so the module
|
||||
can land independently (it is a sibling todo).
|
||||
"""
|
||||
from local_encoder import classify_zero_shot
|
||||
|
||||
if backend == "encoder" and head_off:
|
||||
_disable_trainable_head()
|
||||
|
||||
local_decision = None
|
||||
if backend == "decision":
|
||||
import local_decision
|
||||
|
||||
decision = decision or {}
|
||||
|
||||
# Per (task, noise-level) predictions, keyed enough to build everything.
|
||||
# gold_history[noise] = list of (task_id, gold, predicted, confidence)
|
||||
gold_history: dict[str, list[tuple[str, str, str, float]]] = {
|
||||
@@ -249,6 +288,24 @@ def run_eval(
|
||||
for task_id, gold, _prompt in clean_tasks:
|
||||
level_preds: list[tuple[str, float]] = []
|
||||
for level in noise_levels:
|
||||
if backend == "decision":
|
||||
if decision is None:
|
||||
raise RuntimeError(
|
||||
"--backend decision but decision config is None"
|
||||
)
|
||||
if local_decision is None:
|
||||
raise RuntimeError(
|
||||
"local_decision module not loaded"
|
||||
)
|
||||
predicted, confidence, _coverage = local_decision.classify_category(
|
||||
wrapped[task_id][level],
|
||||
base_url=decision.get("base_url", "http://localhost:11434/api/chat"),
|
||||
model=decision.get("model", "qwen3.5:4b"),
|
||||
num_ctx=decision.get("num_ctx", 8192),
|
||||
timeout_s=decision.get("timeout_s", 120.0),
|
||||
coverage_min=decision.get("coverage_min", 0.0),
|
||||
)
|
||||
else:
|
||||
predicted, confidence = classify_zero_shot(
|
||||
wrapped[task_id][level],
|
||||
categories,
|
||||
@@ -365,7 +422,7 @@ def _print_footer(
|
||||
|
||||
# --- CLI ------------------------------------------------------------------
|
||||
|
||||
def main() -> int:
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
ap = argparse.ArgumentParser(description=__doc__)
|
||||
ap.add_argument("--tasks", default=TASKS_PATH)
|
||||
ap.add_argument(
|
||||
@@ -385,7 +442,46 @@ def main() -> int:
|
||||
default=list(NOISE_LEVELS),
|
||||
help="noise levels to run (default: all)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--backend",
|
||||
choices=["encoder", "decision"],
|
||||
default="encoder",
|
||||
help="classifier backend: encoder (local_encoder) or decision (local_decision)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--head-off",
|
||||
action="store_true",
|
||||
help="encoder: disable the trained logistic-regression head (force zero-shot)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--decision-base-url",
|
||||
help="decision backend Ollama base URL (defaults to classifier.decision.base_url)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--decision-model",
|
||||
help="decision backend model id (defaults to classifier.decision.model)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--decision-num-ctx",
|
||||
type=int,
|
||||
help="decision backend context window (defaults to classifier.decision.num_ctx)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--decision-timeout",
|
||||
type=int,
|
||||
help="decision backend per-call timeout in seconds (defaults to classifier.decision.timeout_s)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--decision-coverage-min",
|
||||
type=float,
|
||||
help="decision backend coverage floor (defaults to classifier.decision.coverage_min)",
|
||||
)
|
||||
ap.add_argument("--dry-run", action="store_true")
|
||||
return ap
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = _build_parser()
|
||||
args = ap.parse_args()
|
||||
|
||||
cfg = load_config("config/config.yaml")
|
||||
@@ -396,6 +492,26 @@ def main() -> int:
|
||||
if args.categories:
|
||||
categories = [c.strip() for c in args.categories.split(",")]
|
||||
|
||||
# Decision-backend config. ``classifier.decision`` is the config block
|
||||
# local_decision's wiring adds; it may not exist yet on this branch, so
|
||||
# read it defensively (getattr) and let the CLI flags override it.
|
||||
decision: Optional[dict] = None
|
||||
if args.backend == "decision":
|
||||
dec = getattr(cfg.classifier, "decision", None)
|
||||
decision = {
|
||||
"base_url": args.decision_base_url
|
||||
or getattr(dec, "base_url", "http://localhost:11434/api/chat"),
|
||||
"model": args.decision_model
|
||||
or getattr(dec, "model", "qwen3.5:4b"),
|
||||
"num_ctx": args.decision_num_ctx
|
||||
or getattr(dec, "num_ctx", 8192),
|
||||
"timeout_s": args.decision_timeout
|
||||
or getattr(dec, "timeout_s", 120),
|
||||
"coverage_min": args.decision_coverage_min
|
||||
if args.decision_coverage_min is not None
|
||||
else getattr(dec, "coverage_min", 0.0),
|
||||
}
|
||||
|
||||
tasks = load_scoreable_tasks(args.tasks)
|
||||
if not tasks:
|
||||
print("no scoreable tasks", file=sys.stderr)
|
||||
@@ -423,13 +539,27 @@ def main() -> int:
|
||||
print(f" {cat:20s} {counts[cat]} tasks")
|
||||
|
||||
if args.dry_run:
|
||||
print("\n--dry-run: no encoder load, no transformers/torch")
|
||||
print(f" model={model_id} device={device}")
|
||||
print("\n--dry-run: no classifier load, no transformers/torch")
|
||||
print(f" backend={args.backend} model={model_id} device={device}")
|
||||
if args.backend == "decision" and decision is not None:
|
||||
print(
|
||||
f" decision model={decision['model']} "
|
||||
f"base_url={decision['base_url']}"
|
||||
)
|
||||
elif args.head_off:
|
||||
print(" trainable head: OFF (--head-off)")
|
||||
print(f" candidate categories ({len(categories)}): {', '.join(categories)}")
|
||||
return 0
|
||||
|
||||
results = run_eval(
|
||||
tasks, categories, model_id=model_id, device=device, noise_levels=noise_levels,
|
||||
tasks,
|
||||
categories,
|
||||
model_id=model_id,
|
||||
device=device,
|
||||
noise_levels=noise_levels,
|
||||
backend=args.backend,
|
||||
head_off=args.head_off,
|
||||
decision=decision,
|
||||
)
|
||||
gold_history = results["gold_history"]
|
||||
task_verdicts = results["task_verdicts"]
|
||||
|
||||
288
src/local_decision.py
Normal file
288
src/local_decision.py
Normal file
@@ -0,0 +1,288 @@
|
||||
"""Parse logprobs from a local Ollama classifier response.
|
||||
|
||||
The local classifier (Ollama) returns a response with top-level ``logprobs`` —
|
||||
an array of position logprob objects, each carrying a ``top_logprobs`` list.
|
||||
This module extracts per-option-letter confidence from those logprobs so the
|
||||
routing layer can decide whether the classifier's choice is trustworthy.
|
||||
|
||||
The classifier prompts the local LLM with four lettered options (A--D). The
|
||||
response looks like::
|
||||
|
||||
{
|
||||
"model": "...",
|
||||
"created_at": "...",
|
||||
"message": { "role": "assistant", "content": "A" },
|
||||
"logprobs": [
|
||||
{
|
||||
"token": "A",
|
||||
"logprob": -0.01,
|
||||
"top_logprobs": [
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "B", "logprob": -3.2}
|
||||
]
|
||||
},
|
||||
{"token": ".", "logprob": -0.5, "top_logprobs": [...]}
|
||||
],
|
||||
"done": true
|
||||
}
|
||||
|
||||
Ollama appends ``.`` (option separator) tokens between letters in some
|
||||
configs; the token key is stripped of its trailing dot before matching.
|
||||
Non-option tokens like ``thinking``, ``assistant``, ``response`` are ignored
|
||||
entirely.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any, Final, Optional
|
||||
|
||||
import requests
|
||||
|
||||
_DEFAULT_COVERAGE_MIN: float = 0.3
|
||||
|
||||
# Decision-category descriptions tuned against the measured confusion
|
||||
# patterns (gate G, see plans/local-decision-classifier-results.md):
|
||||
# 1. coding_general → debugging (code-tracing tasks like "what does this
|
||||
# function return" were read as bug-fixing)
|
||||
# 2. coding_general → reasoning_math (algorithm-implementation tasks were
|
||||
# read as "algorithm analysis")
|
||||
# 3. file_summarization → debugging (code-gotcha summaries were read as
|
||||
# bug-fixing)
|
||||
#
|
||||
# Fixes: widen coding_general to include tracing/implementing, narrow
|
||||
# debugging to explicitly "fixing a bug", and narrow reasoning_math to pure
|
||||
# word problems (no code).
|
||||
#
|
||||
# `tool_use_agentic` is EXCLUDED — the local classifier only presents four
|
||||
# options to the LLM and this slot never appears in the decision prompt.
|
||||
_DECISION_DESCRIPTIONS: dict[str, str] = {
|
||||
"coding_general": (
|
||||
"writing code, implementing a feature or algorithm, or tracing what "
|
||||
"existing code returns"
|
||||
),
|
||||
"coding_refactor": (
|
||||
"refactoring or restructuring existing code without changing its behavior"
|
||||
),
|
||||
"debugging": "fixing a bug in code that is broken or produces wrong output",
|
||||
"docs_writing": (
|
||||
"writing documentation, API docs, or technical explanations (not code)"
|
||||
),
|
||||
"summarization": "summarizing or condensing a long text",
|
||||
"file_summarization": (
|
||||
"summarizing or explaining what a single source file or code module "
|
||||
"contains or does"
|
||||
),
|
||||
"diff_checking": "reviewing a code diff or comparing code changes",
|
||||
"translation": "translating text from one language to another",
|
||||
"reasoning_math": (
|
||||
"solving a pure math or logic word problem (no code to write or run)"
|
||||
),
|
||||
"general_chat": "casual conversation or a general question",
|
||||
}
|
||||
|
||||
_SYSTEM_PROMPT: Final[str] = "answer with the letter only"
|
||||
_ANSWER_SUFFIX: Final[str] = "Answer with the letter only."
|
||||
_NUM_PREDICT: Final[int] = 1
|
||||
_TEMPERATURE: Final[int] = 0
|
||||
_TOP_LOGBPROBS: Final[int] = 20
|
||||
|
||||
|
||||
def parse_logprobs(
|
||||
response_json: dict[str, Any],
|
||||
option_letters: list[str],
|
||||
*,
|
||||
coverage_min: float = _DEFAULT_COVERAGE_MIN,
|
||||
) -> tuple[str, float, float]:
|
||||
"""Extract the classifier's chosen letter and its confidence.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
response_json:
|
||||
The raw Ollama ``/api/chat`` response dict. The logprobs are read
|
||||
from ``response_json["logprobs"]`` (top level).
|
||||
option_letters:
|
||||
The set of valid option labels, e.g. ``["A", "B", "C", "D"]``.
|
||||
coverage_min:
|
||||
Minimum total option-mass to accept the result. Raised as a
|
||||
:exc:`RuntimeError` when the combined ``exp(logprob)`` across all
|
||||
option letters falls below this threshold.
|
||||
|
||||
Returns
|
||||
-------
|
||||
tuple[str, float, float]
|
||||
``(label, confidence, coverage)`` where *confidence* is the winning
|
||||
letter's share of total option mass and *coverage* is the total mass.
|
||||
|
||||
Raises
|
||||
------
|
||||
RuntimeError
|
||||
When ``coverage < coverage_min`` or when no logprobs are found.
|
||||
"""
|
||||
logprobs_list: list[dict[str, Any]] = response_json.get("logprobs", [])
|
||||
|
||||
if not logprobs_list:
|
||||
raise RuntimeError(
|
||||
"parse_logprobs: no logprobs found in classifier response"
|
||||
)
|
||||
|
||||
# Accumulate exp(logprob) mass per option letter.
|
||||
mass: dict[str, float] = {letter: 0.0 for letter in option_letters}
|
||||
|
||||
for position in logprobs_list:
|
||||
for entry in position.get("top_logprobs", []):
|
||||
raw_token: str = entry.get("token", "")
|
||||
# Strip whitespace, then trailing dot (Ollama option separator).
|
||||
token = raw_token.strip().rstrip(".")
|
||||
|
||||
if token not in mass:
|
||||
# Ignore non-option tokens like "thinking", "assistant", "."
|
||||
continue
|
||||
|
||||
logprob: Optional[float] = entry.get("logprob")
|
||||
if logprob is None:
|
||||
continue
|
||||
mass[token] += math.exp(logprob)
|
||||
|
||||
total = sum(mass.values())
|
||||
|
||||
if total < coverage_min:
|
||||
raise RuntimeError(
|
||||
f"parse_logprobs: coverage {total:.4f} below "
|
||||
f"minimum {coverage_min}"
|
||||
)
|
||||
|
||||
# Winner is the option with the most accumulated mass.
|
||||
label = max(mass, key=lambda k: mass[k])
|
||||
confidence = mass[label] / total if total > 0 else 0.0
|
||||
coverage = total
|
||||
|
||||
return (label, confidence, coverage)
|
||||
|
||||
|
||||
def classify_choice(
|
||||
text: str,
|
||||
options: dict[str, str],
|
||||
*,
|
||||
base_url: str,
|
||||
model: str,
|
||||
num_ctx: int,
|
||||
timeout_s: float,
|
||||
coverage_min: float = _DEFAULT_COVERAGE_MIN,
|
||||
) -> tuple[str, float, float]:
|
||||
"""Ask a local Ollama model to choose one of the given options.
|
||||
|
||||
Prompts the model with the task ``text`` and a lettered list of
|
||||
``options`` (e.g. ``{"A": "writing new code", "B": "refactoring"}``),
|
||||
requesting a single-letter answer. The model's first-token logprobs are
|
||||
then parsed into ``(label, confidence, coverage)`` via
|
||||
:func:`parse_logprobs`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
text:
|
||||
The task text to classify.
|
||||
options:
|
||||
Mapping of option letter (e.g. ``"A"``) to a human-readable
|
||||
description of that option.
|
||||
base_url:
|
||||
Base URL of the Ollama server, e.g. ``http://localhost:11434``.
|
||||
model:
|
||||
Name of the Ollama model to call.
|
||||
num_ctx:
|
||||
Context window (in tokens) to use for the request.
|
||||
timeout_s:
|
||||
Timeout in seconds for the HTTP request.
|
||||
coverage_min:
|
||||
Minimum total option-mass to accept the result. Passed through to
|
||||
:func:`parse_logprobs`.
|
||||
|
||||
Returns
|
||||
-------
|
||||
tuple[str, float, float]
|
||||
``(label, confidence, coverage)``.
|
||||
|
||||
Raises
|
||||
------
|
||||
requests.exceptions.RequestException
|
||||
Propagated when the HTTP call fails; callers decide how to cascade.
|
||||
RuntimeError
|
||||
Raised by :func:`parse_logprobs` when coverage is too low or no
|
||||
logprobs were returned.
|
||||
"""
|
||||
option_letters = sorted(options.keys())
|
||||
option_lines = "\n".join(
|
||||
f"{letter}. {description}" for letter, description in options.items()
|
||||
)
|
||||
user_prompt = f"{text}\n\n<task>\n{option_lines}\n{_ANSWER_SUFFIX}"
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": _SYSTEM_PROMPT},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
"options": {
|
||||
"num_predict": _NUM_PREDICT,
|
||||
"temperature": _TEMPERATURE,
|
||||
"num_ctx": num_ctx,
|
||||
},
|
||||
"stream": False,
|
||||
"think": False,
|
||||
"logprobs": True,
|
||||
"top_logprobs": _TOP_LOGBPROBS,
|
||||
}
|
||||
|
||||
resp = requests.post(
|
||||
base_url + "/api/chat",
|
||||
json=payload,
|
||||
timeout=timeout_s,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
|
||||
return parse_logprobs(
|
||||
resp.json(),
|
||||
option_letters=option_letters,
|
||||
coverage_min=coverage_min,
|
||||
)
|
||||
|
||||
|
||||
def classify_category(
|
||||
text: str,
|
||||
*,
|
||||
base_url: str,
|
||||
model: str,
|
||||
num_ctx: int,
|
||||
timeout_s: float,
|
||||
coverage_min: float = _DEFAULT_COVERAGE_MIN,
|
||||
) -> tuple[str, float, float]:
|
||||
"""Ask a local Ollama model to classify the task into a known category.
|
||||
|
||||
Wraps :func:`classify_choice`: assigns letters (A, B, ...) to
|
||||
``_DECISION_DESCRIPTIONS`` in its iteration order, calls
|
||||
``classify_choice()`` with the lettered dict, and maps the winning
|
||||
letter back to the original category name.
|
||||
|
||||
Returns
|
||||
-------
|
||||
tuple[str, float, float]
|
||||
``(category_name, confidence, coverage)``.
|
||||
"""
|
||||
letters = [chr(ord("A") + i) for i in range(len(_DECISION_DESCRIPTIONS))]
|
||||
options = dict(zip(letters, _DECISION_DESCRIPTIONS.values()))
|
||||
letter, confidence, coverage = classify_choice(
|
||||
text,
|
||||
options,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
num_ctx=num_ctx,
|
||||
timeout_s=timeout_s,
|
||||
coverage_min=coverage_min,
|
||||
)
|
||||
# Map the winning letter back to the category name.
|
||||
for i, letter_name in enumerate(letters):
|
||||
if letter_name == letter:
|
||||
category = list(_DECISION_DESCRIPTIONS.keys())[i]
|
||||
return (category, confidence, coverage)
|
||||
# Fallback (should never happen if classification succeeded).
|
||||
return (letter, confidence, coverage)
|
||||
32
tests/fixtures/real_ollama_response.json
vendored
Normal file
32
tests/fixtures/real_ollama_response.json
vendored
Normal file
@@ -0,0 +1,32 @@
|
||||
{
|
||||
"model": "qwen3.5:4b",
|
||||
"created_at": "2026-09-28T12:00:00Z",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": ""
|
||||
},
|
||||
"logprobs": [
|
||||
{
|
||||
"token": "A",
|
||||
"logprob": -0.012,
|
||||
"top_logprobs": [
|
||||
{"token": "A", "logprob": -0.012},
|
||||
{"token": "B", "logprob": -3.84}
|
||||
]
|
||||
},
|
||||
{
|
||||
"token": ".",
|
||||
"logprob": -0.45,
|
||||
"top_logprobs": [
|
||||
{"token": ".", "logprob": -0.45},
|
||||
{"token": ",", "logprob": -1.92}
|
||||
]
|
||||
}
|
||||
],
|
||||
"done": true,
|
||||
"total_duration": 4200000000,
|
||||
"load_duration": 15000000,
|
||||
"prompt_eval_count": 847,
|
||||
"eval_count": 1,
|
||||
"eval_duration": 380000000
|
||||
}
|
||||
@@ -167,6 +167,22 @@ def test_get_reports_no_resolution_when_pinned_not_auto(tmp_path):
|
||||
assert body["cloud_primary"]["model"] == "deepseek-v4-flash"
|
||||
|
||||
|
||||
def test_get_classifier_config_includes_decision_block(tmp_path):
|
||||
"""GET /admin/api/classifier-config includes decision block when mode is local_decision."""
|
||||
client, _config_yaml, _local_yaml = _client(tmp_path)
|
||||
resp = client.post(
|
||||
"/admin/api/classifier-config",
|
||||
json={"mode": "local_decision", "decision": {"coverage_min": 0.42}},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
resp = client.get("/admin/api/classifier-config")
|
||||
body = resp.json()
|
||||
assert body["mode"] == "local_decision"
|
||||
assert body.get("decision") is not None
|
||||
assert body["decision"]["coverage_min"] == 0.42
|
||||
|
||||
|
||||
# --- POST: validated the same way config load is ------------------------
|
||||
|
||||
|
||||
@@ -292,3 +308,26 @@ def test_post_response_names_a_restart_is_required(tmp_path):
|
||||
client, _config_yaml, _local_yaml = _client(tmp_path)
|
||||
resp = client.post("/admin/api/classifier-config", json={"mode": "local_llm"})
|
||||
assert "restart" in resp.json()["message"].lower()
|
||||
|
||||
|
||||
def test_post_local_decision_mode_returns_200(tmp_path):
|
||||
"""local_decision mode with an empty decision block is accepted."""
|
||||
client, _config_yaml, local_yaml = _client(tmp_path)
|
||||
resp = client.post(
|
||||
"/admin/api/classifier-config",
|
||||
json={"mode": "local_decision", "decision": {}},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
written = yaml.safe_load(local_yaml.read_text())
|
||||
assert written["classifier"]["mode"] == "local_decision"
|
||||
assert written["classifier"]["decision"] == {}
|
||||
|
||||
|
||||
def test_post_invalid_mode_returns_422(tmp_path):
|
||||
"""A mode that is not in the Literal is rejected."""
|
||||
client, _config_yaml, _local_yaml = _client(tmp_path)
|
||||
resp = client.post(
|
||||
"/admin/api/classifier-config",
|
||||
json={"mode": "invalid_mode"},
|
||||
)
|
||||
assert resp.status_code == 422
|
||||
|
||||
@@ -97,6 +97,7 @@ def test_config_GET_returns_allowlisted_values(client):
|
||||
"session_cache.enabled",
|
||||
"session_cache.staleness_seconds",
|
||||
"verification.local_llm_enabled",
|
||||
"local_vision.enabled",
|
||||
"objective.incumbent_cache_pricing",
|
||||
"objective.incumbent_challenger_cache_rate",
|
||||
"pinch.enabled",
|
||||
@@ -186,6 +187,42 @@ def test_config_pinch_prefix_probe_is_allowlisted_at_the_right_path(client):
|
||||
assert _CONFIG_ALLOWLIST["pinch.prefix_probe"] == ("pinch", "prefix_probe")
|
||||
|
||||
|
||||
def test_config_local_vision_enabled_round_trips_through_the_real_validator(client):
|
||||
"""local_vision.enabled persists to the temp overlay and reloads as a bool.
|
||||
|
||||
Same reasoning as the ``pinch.prefix_probe`` round-trip: the value that
|
||||
reaches disk has to be one the config loader will take back on the next
|
||||
boot. ``local_vision.enabled`` is read per request in the dispatcher, so a
|
||||
write that lands as the string ``"false"`` would silently keep the fallback
|
||||
on. Assert a real ``False`` survives the merged reload.
|
||||
"""
|
||||
tc, config_yaml = client
|
||||
local_yaml = config_yaml.with_name("config.local.yaml")
|
||||
|
||||
before = tc.get("/admin/api/config").json()
|
||||
assert before["local_vision.enabled"] == {"value": True, "source": "base"}
|
||||
|
||||
resp = tc.post("/admin/api/config/local_vision.enabled", json={"value": False})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["value"] is False
|
||||
|
||||
after = tc.get("/admin/api/config").json()
|
||||
assert after["local_vision.enabled"] == {"value": False, "source": "overlay"}
|
||||
|
||||
# The base file never carries an admin write.
|
||||
assert yaml.safe_load(config_yaml.read_text())["local_vision"]["enabled"] is True
|
||||
assert yaml.safe_load(local_yaml.read_text())["local_vision"]["enabled"] is False
|
||||
|
||||
# And the merged result parses through RouterConfig, which is what the
|
||||
# service does on its next boot.
|
||||
loaded = load_config(str(config_yaml), include_overlay=True)
|
||||
assert loaded.local_vision.enabled is False
|
||||
# The rest of the local_vision block is untouched by the overlay merge.
|
||||
base_vision = yaml.safe_load(config_yaml.read_text())["local_vision"]
|
||||
assert loaded.local_vision.model == base_vision["model"]
|
||||
assert loaded.local_vision.base_url == base_vision["base_url"]
|
||||
|
||||
|
||||
def test_config_incumbent_knobs_round_trip_through_the_real_validator(client):
|
||||
"""Both Wave 2 knobs survive a write as genuine Python types.
|
||||
|
||||
|
||||
@@ -707,3 +707,39 @@ def test_render_board_reads_decision_outcomes_and_render_live_keeps_verdict_mix(
|
||||
assert " ${share}: " not in board
|
||||
live = _function_body(html, "renderLive")
|
||||
assert "verdict_mix" in live
|
||||
|
||||
|
||||
def test_controls_html_has_local_decision_fields():
|
||||
"""controls.html contains all 7 local_decision field data-key attributes.
|
||||
|
||||
The classifierModeFieldsHtml function must render fields for base_url,
|
||||
model, num_ctx, timeout_s, confidence_min, coverage_min, and tier_enabled
|
||||
when mode is local_decision.
|
||||
"""
|
||||
html = (ROOT / "admin" / "frontend" / "controls.html").read_text()
|
||||
assert 'data-key="decision.base_url"' in html
|
||||
assert 'data-key="decision.model"' in html
|
||||
assert 'data-key="decision.num_ctx"' in html
|
||||
assert 'data-key="decision.timeout_s"' in html
|
||||
assert 'data-key="decision.confidence_min"' in html
|
||||
assert 'data-key="decision.coverage_min"' in html
|
||||
assert 'data-key="decision.tier_enabled"' in html
|
||||
|
||||
|
||||
def test_controls_html_local_decision_field_rendering():
|
||||
"""classifierModeFieldsHtml has a local_decision case rendering 7 fields.
|
||||
|
||||
The local_decision block must exist in the function and must reference
|
||||
data.decision for field values.
|
||||
"""
|
||||
html = (ROOT / "admin" / "frontend" / "controls.html").read_text()
|
||||
fields = _function_body(html, "classifierModeFieldsHtml")
|
||||
assert "mode === 'local_decision'" in fields
|
||||
assert "data && data.decision" in fields
|
||||
assert "dec.base_url" in fields
|
||||
assert "dec.model" in fields
|
||||
assert "dec.num_ctx" in fields
|
||||
assert "dec.timeout_s" in fields
|
||||
assert "dec.confidence_min" in fields
|
||||
assert "dec.coverage_min" in fields
|
||||
assert "dec.tier_enabled" in fields
|
||||
|
||||
@@ -46,10 +46,14 @@ against. Two clauses cut it to the knobs an operator would plausibly turn:
|
||||
definitions), ``profiles`` (its own CRUD surface, not a scalar knob),
|
||||
``local_energy`` (machine-specific, lives in the gitignored overlay by
|
||||
design), ``tiers`` / ``tiering`` / ``proficiency`` / ``context`` (catalog and
|
||||
scoring structure), ``local_vision``, and ``classifier`` — the last has its
|
||||
scoring structure), and ``classifier`` — the last has its
|
||||
own dedicated admin card and endpoint pair (``/admin/api/classifier-config``,
|
||||
``/admin/api/cloud-fallback-config``), so measuring it against the GENERIC
|
||||
allowlist would report every field as missing while the card covers it.
|
||||
``local_vision`` is reached now that ``local_vision.enabled`` has a persisted
|
||||
control; the rest of the section (``base_url``, ``api_key_env``, ``model``)
|
||||
is deployment wiring out by clause 2, and the timeout/image limits are
|
||||
excused in ``DELIBERATELY_NOT_IN_ADMIN``.
|
||||
The known limitation: a brand-new section with no control at all is out of
|
||||
scope and unchecked. The first control added under it drags every one of its
|
||||
knobs into scope at once, which is the intended moment to decide.
|
||||
@@ -215,6 +219,21 @@ DELIBERATELY_NOT_IN_ADMIN: dict[str, str] = {
|
||||
"relevance-path shape internal to a covered master switch "
|
||||
"(pinch.relevance.enabled)."
|
||||
),
|
||||
# --- local_vision: the master switch (local_vision.enabled) is covered;
|
||||
# --- the timeout and image limits are tuning for the configured local
|
||||
# --- Ollama vision model, changed with the deployment it points at.
|
||||
"local_vision.timeout_seconds": (
|
||||
"tuning for the configured local vision model; the operator's decision "
|
||||
"is whether the fallback runs at all (local_vision.enabled, covered)."
|
||||
),
|
||||
"local_vision.max_images": (
|
||||
"request-shape limit for the local vision fallback; tuning for the "
|
||||
"configured local model, changed with the deployment."
|
||||
),
|
||||
"local_vision.max_image_bytes": (
|
||||
"request-shape limit for the local vision fallback; tuning for the "
|
||||
"configured local model, changed with the deployment."
|
||||
),
|
||||
# --- measured constants, re-derived rather than dialled ------------------
|
||||
"objective.assumed_cache_rate": (
|
||||
"a measured property of this deployment's traffic (0.917, token-weighted "
|
||||
|
||||
@@ -735,6 +735,24 @@ def test_no_vision_model_uses_local_fallback_when_enabled(router, monkeypatch):
|
||||
), "the image_url part must survive into the local call"
|
||||
|
||||
|
||||
def test_no_vision_model_with_local_fallback_off_is_a_422(router, monkeypatch):
|
||||
"""local_vision.enabled off: no cloud vision candidate 422s, no local call."""
|
||||
client, calls, db_path = router
|
||||
_drop_cheap_by_tier(dispatcher, db_path)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False)
|
||||
_local_vision_fake(monkeypatch, calls, content="must not be used")
|
||||
|
||||
resp = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={"model": "auto", "messages": _image_messages()},
|
||||
)
|
||||
|
||||
assert resp.status_code == 422
|
||||
assert "vision" in resp.json()["detail"]
|
||||
local = [c for c in calls if c.get("local")]
|
||||
assert not local, "the local vision fallback must not run when disabled"
|
||||
|
||||
|
||||
def test_local_fallback_skips_when_json_mode_is_requested(router, monkeypatch):
|
||||
"""Local vision answers in prose; a json_object request must not use it.
|
||||
|
||||
|
||||
@@ -177,3 +177,88 @@ def test_local_encoder_unaffected_by_the_cloud_llm_validator(raw):
|
||||
cfg["classifier"]["mode"] = "local_encoder"
|
||||
cfg["classifier"]["encoder"] = {}
|
||||
RouterConfig(**cfg) # must not raise
|
||||
|
||||
|
||||
# --- local_decision: needs the decision block -----------------------------
|
||||
|
||||
|
||||
def test_local_decision_config_loads_with_defaults(raw):
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
cfg["classifier"]["decision"] = {}
|
||||
loaded = RouterConfig(**cfg)
|
||||
assert loaded.classifier.decision.base_url == "http://localhost:11434"
|
||||
assert loaded.classifier.decision.model == "qwen3.5:4b"
|
||||
assert loaded.classifier.decision.num_ctx == 8192
|
||||
assert loaded.classifier.decision.timeout_s == 10
|
||||
assert loaded.classifier.decision.confidence_min == 0.5
|
||||
assert loaded.classifier.decision.coverage_min == 0.3
|
||||
assert loaded.classifier.decision.tier_enabled is False
|
||||
|
||||
|
||||
def test_local_decision_config_rejects_unknown_fields(raw):
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
cfg["classifier"]["decision"] = {"bogus_field": 1}
|
||||
with pytest.raises(ValidationError, match="extra_forbidden"):
|
||||
RouterConfig(**cfg)
|
||||
|
||||
|
||||
def test_local_decision_config_rejects_invalid_confidence_min(raw):
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
cfg["classifier"]["decision"] = {"confidence_min": 1.5}
|
||||
with pytest.raises(ValueError, match=r"must be in \[0.0, 1.0\]"):
|
||||
RouterConfig(**cfg)
|
||||
|
||||
|
||||
def test_missing_decision_block_rejected_for_local_decision_mode(raw):
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
with pytest.raises(ValueError, match="classifier.decision is not"):
|
||||
RouterConfig(**cfg)
|
||||
|
||||
|
||||
def test_local_decision_is_valid_mode(raw):
|
||||
"""local_decision must be accepted as a classifier.mode literal alongside
|
||||
local_llm/cloud_llm/local_encoder -- the entrypoint for the new mode."""
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
cfg["classifier"]["decision"] = {}
|
||||
loaded = RouterConfig(**cfg)
|
||||
assert loaded.classifier.mode == "local_decision"
|
||||
|
||||
|
||||
# --- local_energy metering uses the right base_url per mode ---------------
|
||||
|
||||
|
||||
def test_local_decision_metering_uses_decision_base_url(raw):
|
||||
"""local_decision mode meters classifier.decision.base_url, not
|
||||
classifier.base_url."""
|
||||
# --- loopback → metered ---
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_decision"
|
||||
cfg["classifier"]["decision"] = {"base_url": "http://localhost:11434"}
|
||||
cfg["local_energy"] = {"enabled": True, "meter": "nvidia_smi", "tariff_usd_per_kwh": 0.10}
|
||||
loaded = RouterConfig(**cfg)
|
||||
assert loaded._sites_cache["classify"] is True
|
||||
|
||||
# --- remote → not metered ---
|
||||
cfg2 = copy.deepcopy(raw)
|
||||
cfg2["classifier"]["mode"] = "local_decision"
|
||||
cfg2["classifier"]["decision"] = {"base_url": "http://remote:11434"}
|
||||
cfg2["local_energy"] = {"enabled": True, "meter": "nvidia_smi", "tariff_usd_per_kwh": 0.10}
|
||||
loaded2 = RouterConfig(**cfg2)
|
||||
assert loaded2._sites_cache["classify"] is False
|
||||
|
||||
|
||||
def test_local_encoder_metering_unconditional(raw):
|
||||
"""local_encoder runs in-process, so classify is always metered
|
||||
regardless of base_url."""
|
||||
cfg = copy.deepcopy(raw)
|
||||
cfg["classifier"]["mode"] = "local_encoder"
|
||||
cfg["classifier"]["encoder"] = {}
|
||||
cfg["classifier"]["base_url"] = "http://remote:11434/v1"
|
||||
cfg["local_energy"] = {"enabled": True, "meter": "nvidia_smi", "tariff_usd_per_kwh": 0.10}
|
||||
loaded = RouterConfig(**cfg)
|
||||
assert loaded._sites_cache["classify"] is True
|
||||
|
||||
@@ -12,12 +12,18 @@ signal. An intentionally configured cloud primary is not degraded.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import time
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
import dispatcher
|
||||
import local_decision
|
||||
import local_energy
|
||||
import local_encoder
|
||||
import session_cache
|
||||
|
||||
@@ -336,6 +342,104 @@ def test_local_encoder_below_threshold_uses_the_real_cascade(monkeypatch):
|
||||
assert got.task_category == "coding_refactor"
|
||||
|
||||
|
||||
# --- mode: local_decision ---------------------------------------------------
|
||||
|
||||
|
||||
def _decision_conf(confidence_min=0.5, coverage_min=0.3, tier_enabled=False):
|
||||
return SimpleNamespace(
|
||||
base_url="http://localhost:11434",
|
||||
model="stub-model",
|
||||
num_ctx=8192,
|
||||
timeout_s=10,
|
||||
confidence_min=confidence_min,
|
||||
coverage_min=coverage_min,
|
||||
tier_enabled=tier_enabled,
|
||||
)
|
||||
|
||||
|
||||
def _logprob_response(label="coding_refactor"):
|
||||
return {
|
||||
"logprobs": [
|
||||
{
|
||||
"token": label,
|
||||
"logprob": -0.01,
|
||||
"top_logprobs": [{"token": label, "logprob": -0.01}],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_local_decision_success_records_source_classifier(monkeypatch):
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf(confidence_min=0.5)
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_category",
|
||||
lambda *a, **k: calls.append((a, k))
|
||||
or ("coding_refactor", 0.83, 0.9),
|
||||
)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this function")
|
||||
|
||||
assert got.source == "classifier"
|
||||
assert got.task_category == "coding_refactor"
|
||||
assert got.confidence == pytest.approx(0.83)
|
||||
assert calls[0][0][0] == "refactor this function"
|
||||
assert calls[0][1]["base_url"] == "http://localhost:11434"
|
||||
|
||||
|
||||
def test_local_decision_uses_fallback_tier(monkeypatch):
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 3)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf(confidence_min=0.5)
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_category",
|
||||
lambda *a, **k: ("general_chat", 0.9, 0.9),
|
||||
)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "hello")
|
||||
assert got.task_tier == 3
|
||||
|
||||
|
||||
def test_local_decision_below_confidence_cascades(monkeypatch):
|
||||
"""A below-confidence local_decision verdict is a failure, not a guess."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf(confidence_min=0.6)
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_category",
|
||||
lambda *a, **k: ("general_chat", 0.2, 0.9),
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="confidence"):
|
||||
dispatcher._classify_via_local_decision("sys", "do a thing")
|
||||
|
||||
|
||||
def test_local_decision_skip_reason_respected(monkeypatch):
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf(confidence_min=0.5)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
dispatcher, "_local_classifier_skip_reason", lambda: "gaming_mode"
|
||||
)
|
||||
|
||||
with pytest.raises(dispatcher._ClassifierSkipped, match="gaming_mode"):
|
||||
dispatcher._classify_via_local_decision("sys", "do a thing")
|
||||
|
||||
|
||||
# --- startup: local_encoder mode must fail loudly, not on the first request -
|
||||
|
||||
|
||||
@@ -385,3 +489,527 @@ def test_startup_check_surfaces_a_missing_dependency_loudly(monkeypatch):
|
||||
|
||||
with pytest.raises(ImportError, match="pip install -r requirements-encoder.txt"):
|
||||
dispatcher._ensure_classifier_mode_ready()
|
||||
|
||||
|
||||
# --- startup: local_decision mode must warn, not fail, if Ollama is down ----
|
||||
|
||||
def _decision_startup_conf():
|
||||
return _decision_conf(confidence_min=0.5, coverage_min=0.3)
|
||||
|
||||
|
||||
def test_startup_check_validates_decision_mode(monkeypatch):
|
||||
"""A healthy local_decision mode makes one classify_category call
|
||||
and passes."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "decision", _decision_startup_conf())
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_category",
|
||||
lambda *a, **k: calls.append((a, k)) or ("coding_general", 0.9, 0.9),
|
||||
)
|
||||
dispatcher._ensure_classifier_mode_ready() # must not raise
|
||||
assert len(calls) == 1
|
||||
assert calls[0][0][0] == "test task"
|
||||
assert calls[0][1]["base_url"] == "http://localhost:11434"
|
||||
assert calls[0][1]["model"] == "stub-model"
|
||||
|
||||
|
||||
def test_startup_check_warns_on_decision_failure(monkeypatch, caplog):
|
||||
"""Ollama down at boot: warn and continue — the cascade handles failures."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "decision", _decision_startup_conf())
|
||||
|
||||
def fake_classify_category(*a, **k):
|
||||
raise ConnectionError("Ollama is not running")
|
||||
|
||||
monkeypatch.setattr(local_decision, "classify_category", fake_classify_category)
|
||||
|
||||
caplog.set_level("WARNING", logger="llm_router")
|
||||
dispatcher._ensure_classifier_mode_ready() # must NOT raise
|
||||
|
||||
assert any(
|
||||
"classifier_mode_unavailable" in r.getMessage()
|
||||
and "mode=local_decision" in r.getMessage()
|
||||
and "Ollama is not running" in r.getMessage()
|
||||
for r in caplog.records
|
||||
), "a local_decision startup failure must be logged as classifier_mode_unavailable"
|
||||
|
||||
|
||||
def test_startup_check_is_noop_for_non_decision_modes(monkeypatch):
|
||||
"""Encoder mode still fires ensure_available; decision path is untouched."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier,
|
||||
"encoder",
|
||||
SimpleNamespace(model="stub-model", device="cpu", confidence_min=0.5),
|
||||
)
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
local_encoder, "ensure_available", lambda model, device: calls.append((model, device))
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_choice",
|
||||
lambda *a, **k: pytest.fail("classify_choice called for a non-decision mode"),
|
||||
)
|
||||
dispatcher._ensure_classifier_mode_ready()
|
||||
assert calls == [("stub-model", "cpu")]
|
||||
|
||||
|
||||
def test_startup_check_skipped_when_local_compute_off(monkeypatch):
|
||||
"""When local_compute.enabled is False (gaming mode), the local_decision
|
||||
startup probe must not call Ollama — it would load the 4b model onto the GPU.
|
||||
"""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "decision", _decision_startup_conf())
|
||||
monkeypatch.setattr(dispatcher.cfg.local_compute, "enabled", False)
|
||||
|
||||
def probe_should_not_be_called(*a, **k):
|
||||
pytest.fail("classify_category must NOT be called when local_compute.enabled is False")
|
||||
|
||||
monkeypatch.setattr(local_decision, "classify_category", probe_should_not_be_called)
|
||||
|
||||
dispatcher._ensure_classifier_mode_ready()
|
||||
|
||||
|
||||
# --- tier classification in local_decision mode ------------------------------
|
||||
|
||||
|
||||
def test_local_decision_tier_enabled_false_no_tier_calls(monkeypatch):
|
||||
"""When tier_enabled is False (the default) classify_category fires
|
||||
only once — for category — and tier equals fallback_tier."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf(confidence_min=0.5)
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
local_decision,
|
||||
"classify_category",
|
||||
lambda *a, **k: calls.append(1) or ("coding_refactor", 0.83, 0.9),
|
||||
)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this")
|
||||
|
||||
assert got.task_tier == 2
|
||||
assert len(calls) == 1, "classify_category should fire only once for category"
|
||||
|
||||
|
||||
def test_local_decision_tier_enabled_true_makes_tier_call(monkeypatch):
|
||||
"""When tier_enabled is True classify_category fires once for
|
||||
category, and classify_choice fires once for tier via the tier
|
||||
function, and the returned tier matches the tier prompt's
|
||||
option letter mapping."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 1)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=True),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
|
||||
call_count = 0
|
||||
|
||||
def fake_classify_category(*a, **k):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return ("coding_refactor", 0.83, 0.9)
|
||||
|
||||
def fake_classify_choice(*a, **k):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
# Return B → tier 2
|
||||
return ("B", 0.75, 0.85)
|
||||
|
||||
monkeypatch.setattr(local_decision, "classify_category", fake_classify_category)
|
||||
monkeypatch.setattr(local_decision, "classify_choice", fake_classify_choice)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this")
|
||||
|
||||
assert got.task_tier == 2
|
||||
assert got.task_category == "coding_refactor"
|
||||
assert call_count == 2, (
|
||||
"classify_category once for category + classify_choice once for tier"
|
||||
)
|
||||
|
||||
|
||||
def test_local_decision_tier_failure_falls_to_fallback(monkeypatch):
|
||||
"""If the tier call raises (network/model error) tier falls to
|
||||
fallback_tier with zero disruption to the category result."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 3)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=True),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
|
||||
call_count = 0
|
||||
|
||||
def fake_classify_category(*a, **k):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return ("coding_refactor", 0.83, 0.9)
|
||||
|
||||
def fake_classify_choice(*a, **k):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
# Second call: simulate network/model failure
|
||||
raise ConnectionError("Ollama is not running")
|
||||
|
||||
monkeypatch.setattr(local_decision, "classify_category", fake_classify_category)
|
||||
monkeypatch.setattr(local_decision, "classify_choice", fake_classify_choice)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this")
|
||||
|
||||
assert got.task_tier == 3, "tier should fall to fallback on failure"
|
||||
assert got.task_category == "coding_refactor"
|
||||
assert call_count == 2
|
||||
|
||||
|
||||
def test_local_decision_tier_runs_in_parallel_with_category(monkeypatch):
|
||||
"""Category and tier calls execute concurrently when tier_enabled is
|
||||
True. Both must hit a barrier before either completes — proving they
|
||||
run in separate threads inside a shared ThreadPoolExecutor."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=True),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
barrier = threading.Barrier(2, timeout=2)
|
||||
|
||||
def fake_classify_category(*a, **k):
|
||||
time.sleep(0.05)
|
||||
barrier.wait()
|
||||
return ("coding_refactor", 0.83, 0.9)
|
||||
|
||||
def fake_classify_choice(*a, **k):
|
||||
time.sleep(0.05)
|
||||
barrier.wait()
|
||||
return ("B", 0.75, 0.85)
|
||||
|
||||
monkeypatch.setattr(local_decision, "classify_category", fake_classify_category)
|
||||
monkeypatch.setattr(local_decision, "classify_choice", fake_classify_choice)
|
||||
|
||||
start = time.time()
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this")
|
||||
elapsed = time.time() - start
|
||||
|
||||
assert got.task_category == "coding_refactor"
|
||||
assert got.task_tier == 2
|
||||
# With two 50 ms sleeps and a barrier, sequential code takes ≥0.2 s,
|
||||
# parallel code ≈0.05 s. Give generous slack but fail on sequential.
|
||||
assert elapsed < 0.25, f"took {elapsed:.3f}s — calls ran sequentially, not in parallel"
|
||||
|
||||
|
||||
def test_local_decision_metered_logs_correct_duration_power(monkeypatch):
|
||||
def fake_nvidia(*a, **k):
|
||||
return 100.0
|
||||
|
||||
def fake_post(*a, **k):
|
||||
time.sleep(0.2)
|
||||
return SimpleNamespace(
|
||||
json=lambda: {
|
||||
"model": "qwen3.5:4b",
|
||||
"message": {"role": "assistant", "content": "B"},
|
||||
"logprobs": [{
|
||||
"top_logprobs": [
|
||||
{"token": "B", "logprob": -0.1},
|
||||
{"token": "A", "logprob": -2.0},
|
||||
{"token": "C", "logprob": -3.0},
|
||||
]
|
||||
}],
|
||||
},
|
||||
raise_for_status=lambda: None,
|
||||
)
|
||||
|
||||
def capture_log(*a, measurement=None, **k):
|
||||
measurement_sent.append(SimpleNamespace(
|
||||
duration_seconds=measurement.duration_seconds,
|
||||
avg_power_watts=measurement.avg_power_watts,
|
||||
))
|
||||
|
||||
measurement_sent = []
|
||||
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=False),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_energy, "enabled", True)
|
||||
monkeypatch.setattr(
|
||||
dispatcher, "cfg",
|
||||
SimpleNamespace(
|
||||
classifier=dispatcher.cfg.classifier,
|
||||
local_energy=SimpleNamespace(
|
||||
enabled=True,
|
||||
sample_interval_seconds=0.01,
|
||||
call_sites={"classify": True},
|
||||
tariff=0.08,
|
||||
),
|
||||
local_energy_call_sites={"classify": True},
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
monkeypatch.setattr(dispatcher.local_energy, "sample_nvidia_smi", fake_nvidia)
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
monkeypatch.setattr(dispatcher, "_log_local_energy", capture_log)
|
||||
|
||||
got = dispatcher._classify_via_local_decision("sys", "refactor this")
|
||||
assert got.task_category == "coding_refactor"
|
||||
|
||||
assert len(measurement_sent) == 1, "expected one _log_local_energy call"
|
||||
sent = measurement_sent[0]
|
||||
assert sent.duration_seconds >= 0.15, \
|
||||
f"expected >= 0.15, got {sent.duration_seconds}"
|
||||
assert sent.avg_power_watts == 100.0, \
|
||||
f"expected 100.0, got {sent.avg_power_watts}"
|
||||
|
||||
|
||||
# --- HTTP-level tests: letter→category mapping via mocked requests.post ------
|
||||
# These tests fake local_decision.requests.post (the lowest HTTP boundary)
|
||||
# so they exercise classify_choice → parse_logprobs → letter→category mapping
|
||||
# end-to-end, NOT just the return value of a stubbed classify_category.
|
||||
|
||||
|
||||
def make_response(letter="B", top_logprobs=None):
|
||||
"""Create a mock response object like requests.post() returns.
|
||||
|
||||
``top_logprobs`` defaults to three options with B as clear winner.
|
||||
"""
|
||||
if top_logprobs is None:
|
||||
top_logprobs = [
|
||||
{"token": letter, "logprob": -0.1},
|
||||
{"token": "A", "logprob": -2.0},
|
||||
{"token": "C", "logprob": -3.0},
|
||||
]
|
||||
return SimpleNamespace(
|
||||
json=lambda: {
|
||||
"model": "qwen3.5:4b",
|
||||
"message": {"role": "assistant", "content": letter},
|
||||
"logprobs": [{"top_logprobs": top_logprobs}],
|
||||
},
|
||||
raise_for_status=lambda: None,
|
||||
)
|
||||
|
||||
|
||||
def test_classify_category_maps_letter_to_category(monkeypatch):
|
||||
"""When requests.post returns top token 'B', classify_category maps it
|
||||
to the second key of _DECISION_DESCRIPTIONS ('coding_refactor')."""
|
||||
from src.local_decision import _DECISION_DESCRIPTIONS, classify_category
|
||||
|
||||
second_key = list(_DECISION_DESCRIPTIONS.keys())[1]
|
||||
|
||||
captured_calls = []
|
||||
|
||||
def fake_post(*a, **k):
|
||||
captured_calls.append((a, k))
|
||||
return make_response(letter="B")
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
result = classify_category(
|
||||
"test task",
|
||||
base_url="http://localhost:11434",
|
||||
model="qwen3.5:4b",
|
||||
num_ctx=8192,
|
||||
timeout_s=10,
|
||||
coverage_min=0.3,
|
||||
)
|
||||
|
||||
label, confidence, coverage = result
|
||||
assert label == second_key, f"Expected {second_key!r}, got {label!r}"
|
||||
assert len(captured_calls) == 1
|
||||
assert 0.0 < confidence <= 1.0
|
||||
assert coverage > 0.3
|
||||
|
||||
|
||||
def test_dispatcher_returns_category_from_letter(monkeypatch):
|
||||
"""When the model returns letter 'B', _classify_via_local_decision
|
||||
returns the correct category ('coding_refactor') with source='classifier',
|
||||
proving the full requests.post → classify_choice → letter→category mapping
|
||||
works through the dispatcher."""
|
||||
from src.local_decision import _DECISION_DESCRIPTIONS
|
||||
|
||||
second_key = list(_DECISION_DESCRIPTIONS.keys())[1]
|
||||
|
||||
captured_calls = []
|
||||
|
||||
def fake_post(*a, **k):
|
||||
captured_calls.append((a, k))
|
||||
return make_response(letter="B")
|
||||
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=False),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
result = dispatcher._classify_via_local_decision("sys", "test task")
|
||||
|
||||
assert result.task_category == second_key, \
|
||||
f"Expected {second_key!r}, got {result.task_category!r}"
|
||||
assert result.source == "classifier"
|
||||
assert len(captured_calls) == 1
|
||||
|
||||
|
||||
def test_tier_enabled_false_is_default(monkeypatch):
|
||||
"""tier_enabled defaults to False so the normal local_decision path
|
||||
never fires extra HTTP calls without explicit config."""
|
||||
from config import LocalDecisionConfig
|
||||
|
||||
assert LocalDecisionConfig().tier_enabled is False
|
||||
|
||||
|
||||
def test_unreachable_ollama_walks_the_cascade(monkeypatch):
|
||||
"""When local_decision Ollama is unreachable, classify() returns a
|
||||
degraded fallback Classification instead of raising (c06bc36)."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision", _decision_conf()
|
||||
)
|
||||
dispatcher.cfg.classifier.decision.base_url = "http://127.0.0.1:1"
|
||||
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
|
||||
monkeypatch.setattr(
|
||||
local_decision.requests, "post", lambda *a, **k: (_ for _ in ()).throw(
|
||||
requests.ConnectionError("Connection refused")
|
||||
)
|
||||
)
|
||||
|
||||
got = dispatcher.classify("refactor this function", None)
|
||||
|
||||
assert got.source == "fallback"
|
||||
assert got.confidence == 0.0
|
||||
assert got.required_context_tokens == 0
|
||||
assert got.task_tier == dispatcher.cfg.classifier.fallback_tier
|
||||
assert got.task_category == dispatcher.cfg.classifier.fallback_category
|
||||
|
||||
|
||||
def test_local_decision_connection_error_opens_circuit_unmetered(monkeypatch):
|
||||
"""Transport failure in unmetered local_decision opens the classifier
|
||||
backoff and blocks subsequent calls until the circuit closes."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=False),
|
||||
)
|
||||
# Make metering off so we hit the unmetered path
|
||||
monkeypatch.setattr(dispatcher.cfg.local_energy, "enabled", False)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_compute, "enabled", True)
|
||||
|
||||
call_count = [0]
|
||||
def fake_post(*a, **k):
|
||||
call_count[0] += 1
|
||||
raise requests.ConnectionError("Connection refused")
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
# First call — should record failure and return fallback
|
||||
got1 = dispatcher.classify("refactor this function", None)
|
||||
assert got1.source == "fallback"
|
||||
assert dispatcher._last_classifier_failure > 0.0
|
||||
assert dispatcher._classifier_backoff_active() is True
|
||||
|
||||
# Second call — should be SKIPPED (post called exactly once total)
|
||||
got2 = dispatcher.classify("another task", None)
|
||||
assert call_count[0] == 1, f"post called {call_count[0]} times, expected 1"
|
||||
assert got2.source == "fallback"
|
||||
|
||||
# _last_classifier_failure unchanged by second call
|
||||
assert dispatcher._last_classifier_failure > 0.0 # still set from first call
|
||||
|
||||
|
||||
def test_local_decision_connection_error_opens_circuit_metered(monkeypatch):
|
||||
"""Transport failure in metered local_decision opens the classifier
|
||||
backoff and blocks subsequent calls until the circuit closes."""
|
||||
call_count = [0]
|
||||
|
||||
def fake_nvidia(*a, **k):
|
||||
return 100.0
|
||||
|
||||
def fake_post(*a, **k):
|
||||
call_count[0] += 1
|
||||
raise requests.ConnectionError("Connection refused")
|
||||
|
||||
def capture_log(*a, measurement=None, **k):
|
||||
pass # just swallow — we don't care about the measurement
|
||||
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=False),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_energy, "enabled", True)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_energy, "sample_interval_seconds", 0.01)
|
||||
monkeypatch.setattr(dispatcher.cfg, "_sites_cache", {"classify": True})
|
||||
monkeypatch.setattr(dispatcher.cfg.local_compute, "enabled", True)
|
||||
monkeypatch.setattr(dispatcher.local_energy, "sample_nvidia_smi", fake_nvidia)
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
monkeypatch.setattr(dispatcher, "_log_local_energy", capture_log)
|
||||
|
||||
# First call — should record failure and return fallback
|
||||
got1 = dispatcher.classify("refactor this function", None)
|
||||
assert got1.source == "fallback"
|
||||
assert dispatcher._last_classifier_failure > 0.0
|
||||
assert dispatcher._classifier_backoff_active() is True
|
||||
|
||||
# Second call — should be SKIPPED (post called exactly once total)
|
||||
got2 = dispatcher.classify("another task", None)
|
||||
assert call_count[0] == 1, f"post called {call_count[0]} times, expected 1"
|
||||
assert got2.source == "fallback"
|
||||
|
||||
|
||||
def test_local_decision_low_confidence_does_not_open_circuit(monkeypatch):
|
||||
"""A successful classify with low confidence is a real answer, not a
|
||||
transport failure — the circuit must stay closed."""
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
|
||||
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
|
||||
monkeypatch.setattr(
|
||||
dispatcher.cfg.classifier, "decision",
|
||||
_decision_conf(confidence_min=0.5, tier_enabled=False),
|
||||
)
|
||||
monkeypatch.setattr(dispatcher.cfg.local_compute, "enabled", True)
|
||||
|
||||
def fake_post(*a, **k):
|
||||
# Return a response where the winning token has very low logprob
|
||||
# → confidence below 0.5 → below confidence_min → RuntimeError
|
||||
return SimpleNamespace(
|
||||
json=lambda: {
|
||||
"model": "qwen3.5:4b",
|
||||
"message": {"role": "assistant", "content": "Z"},
|
||||
"logprobs": [{
|
||||
"top_logprobs": [
|
||||
{"token": "Z", "logprob": -5.0}, # very low confidence
|
||||
{"token": "A", "logprob": -5.1},
|
||||
{"token": "B", "logprob": -5.2},
|
||||
]
|
||||
}],
|
||||
},
|
||||
raise_for_status=lambda: None,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
# classify() should cascade to fallback because confidence < confidence_min
|
||||
got = dispatcher.classify("refactor this function", None)
|
||||
assert got.source == "fallback" # cascaded
|
||||
|
||||
# But _last_classifier_failure must be 0.0 — this was NOT a transport error
|
||||
assert dispatcher._last_classifier_failure == 0.0
|
||||
assert dispatcher._classifier_backoff_active() is False
|
||||
|
||||
181
tests/test_eval_classifier.py
Normal file
181
tests/test_eval_classifier.py
Normal file
@@ -0,0 +1,181 @@
|
||||
"""Tests for the eval_classifier CLI backends (Todo 18).
|
||||
|
||||
Covers the ``--backend encoder|decision`` and ``--head-off`` flags:
|
||||
flag parsing, the ``--head-off`` effect on the encoder's trainable-head
|
||||
globals, and the decision backend dispatching through
|
||||
``local_decision.classify_choice`` over ``_DECISION_DESCRIPTIONS``.
|
||||
|
||||
``local_decision`` does not exist on this branch yet (it is a sibling todo),
|
||||
so the decision-backend tests stub the module in ``sys.modules``. Nothing here
|
||||
touches the network, transformers or torch.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
import local_encoder
|
||||
from eval_classifier import (
|
||||
_build_parser,
|
||||
_disable_trainable_head,
|
||||
run_eval,
|
||||
)
|
||||
|
||||
|
||||
def _sample_tasks():
|
||||
return [
|
||||
{"id": "t1", "category": "coding_general", "prompt": "write a function"},
|
||||
{"id": "t2", "category": "debugging", "prompt": "fix this crash"},
|
||||
]
|
||||
|
||||
|
||||
CATS = ["coding_general", "debugging"]
|
||||
|
||||
|
||||
# --- flag parsing ---------------------------------------------------------
|
||||
|
||||
def test_backend_defaults_to_encoder():
|
||||
args = _build_parser().parse_args([])
|
||||
assert args.backend == "encoder"
|
||||
assert args.head_off is False
|
||||
|
||||
|
||||
def test_backend_accepts_decision():
|
||||
args = _build_parser().parse_args(["--backend", "decision"])
|
||||
assert args.backend == "decision"
|
||||
|
||||
|
||||
def test_backend_rejects_unknown_choice():
|
||||
with pytest.raises(SystemExit):
|
||||
_build_parser().parse_args(["--backend", "bogus"])
|
||||
|
||||
|
||||
def test_head_off_flag_is_store_true():
|
||||
args = _build_parser().parse_args(["--head-off"])
|
||||
assert args.head_off is True
|
||||
|
||||
|
||||
def test_decision_override_flags_parse():
|
||||
args = _build_parser().parse_args([
|
||||
"--backend", "decision",
|
||||
"--decision-base-url", "http://x:11434/api/chat",
|
||||
"--decision-model", "qwen3.5:4b",
|
||||
"--decision-num-ctx", "8192",
|
||||
"--decision-timeout", "60",
|
||||
"--decision-coverage-min", "0.5",
|
||||
])
|
||||
assert args.decision_base_url == "http://x:11434/api/chat"
|
||||
assert args.decision_model == "qwen3.5:4b"
|
||||
assert args.decision_num_ctx == 8192
|
||||
assert args.decision_timeout == 60
|
||||
assert args.decision_coverage_min == 0.5
|
||||
|
||||
|
||||
# --- --head-off -----------------------------------------------------------
|
||||
|
||||
def test_disable_trainable_head_sets_globals():
|
||||
local_encoder._TRAINABLE_HEAD = object()
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
try:
|
||||
_disable_trainable_head()
|
||||
assert local_encoder._TRAINABLE_HEAD is None
|
||||
assert local_encoder._TRAINABLE_HEAD_RESOLVED is True
|
||||
finally:
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
|
||||
|
||||
def test_run_eval_encoder_head_off_disables_head():
|
||||
# --head-off with the encoder backend must set the globals before scoring.
|
||||
local_encoder._TRAINABLE_HEAD = object()
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
try:
|
||||
with mock.patch(
|
||||
"local_encoder.classify_zero_shot",
|
||||
return_value=("coding_general", 0.9),
|
||||
):
|
||||
run_eval(
|
||||
_sample_tasks(), CATS,
|
||||
model_id="x", device="cpu", noise_levels=("clean",),
|
||||
backend="encoder", head_off=True,
|
||||
)
|
||||
assert local_encoder._TRAINABLE_HEAD is None
|
||||
assert local_encoder._TRAINABLE_HEAD_RESOLVED is True
|
||||
finally:
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
|
||||
|
||||
def test_run_eval_encoder_default_does_not_touch_head():
|
||||
# Without --head-off the encoder backend leaves the globals alone.
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
try:
|
||||
with mock.patch(
|
||||
"local_encoder.classify_zero_shot",
|
||||
return_value=("coding_general", 0.9),
|
||||
):
|
||||
run_eval(
|
||||
_sample_tasks(), CATS,
|
||||
model_id="x", device="cpu", noise_levels=("clean",),
|
||||
backend="encoder", head_off=False,
|
||||
)
|
||||
assert local_encoder._TRAINABLE_HEAD_RESOLVED is False
|
||||
finally:
|
||||
local_encoder._TRAINABLE_HEAD = None
|
||||
local_encoder._TRAINABLE_HEAD_RESOLVED = False
|
||||
|
||||
|
||||
# --- decision backend -----------------------------------------------------
|
||||
|
||||
def _fake_decision_module():
|
||||
mod = mock.Mock()
|
||||
mod._DECISION_DESCRIPTIONS = {
|
||||
"coding_general": "writing or editing code",
|
||||
"debugging": "finding and fixing bugs",
|
||||
}
|
||||
mod.classify_category.return_value = ("coding_general", 0.95, 1.0)
|
||||
return mod
|
||||
|
||||
|
||||
def test_run_eval_decision_uses_classify_category():
|
||||
fake = _fake_decision_module()
|
||||
decision = {
|
||||
"base_url": "http://x:11434/api/chat",
|
||||
"model": "qwen3.5:4b",
|
||||
"num_ctx": 8192,
|
||||
"timeout_s": 60,
|
||||
"coverage_min": 0.0,
|
||||
}
|
||||
with mock.patch.dict(sys.modules, {"local_decision": fake}):
|
||||
results = run_eval(
|
||||
_sample_tasks(), CATS,
|
||||
model_id="x", device="cpu", noise_levels=("clean",),
|
||||
backend="decision", decision=decision,
|
||||
)
|
||||
|
||||
assert fake.classify_category.called
|
||||
_, kwargs = fake.classify_category.call_args
|
||||
assert kwargs["base_url"] == decision["base_url"]
|
||||
assert kwargs["model"] == decision["model"]
|
||||
assert kwargs["num_ctx"] == decision["num_ctx"]
|
||||
assert kwargs["timeout_s"] == decision["timeout_s"]
|
||||
assert kwargs["coverage_min"] == decision["coverage_min"]
|
||||
|
||||
# The returned category name is used directly as the predicted verdict.
|
||||
tv = results["task_verdicts"]
|
||||
assert tv["t1"][1] == "coding_general"
|
||||
assert tv["t2"][1] == "coding_general"
|
||||
|
||||
|
||||
def test_run_eval_decision_defaults_coverage_min():
|
||||
fake = _fake_decision_module()
|
||||
with mock.patch.dict(sys.modules, {"local_decision": fake}):
|
||||
run_eval(
|
||||
_sample_tasks(), CATS,
|
||||
model_id="x", device="cpu", noise_levels=("clean",),
|
||||
backend="decision", decision={"base_url": "u", "model": "m"},
|
||||
)
|
||||
_, kwargs = fake.classify_category.call_args
|
||||
assert kwargs["coverage_min"] == 0.0
|
||||
62
tests/test_eval_heldout.py
Normal file
62
tests/test_eval_heldout.py
Normal file
@@ -0,0 +1,62 @@
|
||||
"""Tests that evals/heldout.yaml is a loadable, scoreable task set.
|
||||
|
||||
The held-out set is a fixed benchmark: 30 tasks across 10 categories, with
|
||||
tool-use tasks excluded from the scoreable set (as in the plan). These tests
|
||||
guard the loader contract against a malformed or drifted file.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
|
||||
|
||||
from eval_classifier import EXCLUDED_CATEGORIES, load_scoreable_tasks
|
||||
|
||||
HELDOUT_PATH = Path(__file__).resolve().parent.parent / "evals" / "heldout.yaml"
|
||||
|
||||
# The 10 categories in the held-out set, one per 3-task group.
|
||||
EXPECTED_CATEGORIES = [
|
||||
"coding_general",
|
||||
"coding_refactor",
|
||||
"debugging",
|
||||
"docs_writing",
|
||||
"summarization",
|
||||
"file_summarization",
|
||||
"diff_checking",
|
||||
"translation",
|
||||
"reasoning_math",
|
||||
"general_chat",
|
||||
]
|
||||
|
||||
|
||||
def test_heldout_file_exists():
|
||||
assert HELDOUT_PATH.exists(), "evals/heldout.yaml must be committed alongside the loader"
|
||||
|
||||
|
||||
def test_heldout_has_30_scoreable_tasks():
|
||||
tasks = load_scoreable_tasks(str(HELDOUT_PATH))
|
||||
assert len(tasks) == 30
|
||||
|
||||
|
||||
def test_heldout_categories_match_expected():
|
||||
tasks = load_scoreable_tasks(str(HELDOUT_PATH))
|
||||
categories = [t["category"] for t in tasks]
|
||||
# The held-out set is grouped: three tasks per category, in order.
|
||||
expected = [c for c in EXPECTED_CATEGORIES for _ in range(3)]
|
||||
assert categories == expected
|
||||
|
||||
|
||||
def test_heldout_tasks_have_id_and_prompt():
|
||||
tasks = load_scoreable_tasks(str(HELDOUT_PATH))
|
||||
for t in tasks:
|
||||
assert t["id"]
|
||||
assert t["prompt"]
|
||||
|
||||
|
||||
def test_heldout_excludes_tool_use():
|
||||
# The scoreable set must never admit tool-use tasks (per the plan).
|
||||
tasks = load_scoreable_tasks(str(HELDOUT_PATH))
|
||||
cats = {t["category"] for t in tasks}
|
||||
assert not (cats & set(EXCLUDED_CATEGORIES))
|
||||
161
tests/test_local_decision_descriptions.py
Normal file
161
tests/test_local_decision_descriptions.py
Normal file
@@ -0,0 +1,161 @@
|
||||
"""Structural tests for local_decision._DECISION_DESCRIPTIONS.
|
||||
|
||||
Verifies:
|
||||
- All 10 expected category keys are present.
|
||||
- `tool_use_agentic` is excluded.
|
||||
- The dict is NOT a shallow copy of local_encoder._CATEGORY_DESCRIPTIONS.
|
||||
- local_encoder.py is NOT modified by this change (git diff is empty).
|
||||
- classify_choice would work with these descriptions (they're valid option values).
|
||||
"""
|
||||
|
||||
import pathlib
|
||||
import subprocess
|
||||
|
||||
|
||||
def _encoder_descriptions():
|
||||
"""Import and return local_encoder._CATEGORY_DESCRIPTIONS."""
|
||||
import local_encoder
|
||||
return local_encoder._CATEGORY_DESCRIPTIONS
|
||||
|
||||
|
||||
def _decision_descriptions():
|
||||
"""Import and return local_decision._DECISION_DESCRIPTIONS."""
|
||||
import local_decision
|
||||
return local_decision._DECISION_DESCRIPTIONS
|
||||
|
||||
|
||||
# Expected keys: all 11 minus tool_use_agentic = 10
|
||||
_EXPECTED_KEYS = sorted([
|
||||
"coding_general",
|
||||
"coding_refactor",
|
||||
"debugging",
|
||||
"docs_writing",
|
||||
"summarization",
|
||||
"file_summarization",
|
||||
"diff_checking",
|
||||
"translation",
|
||||
"reasoning_math",
|
||||
"general_chat",
|
||||
])
|
||||
|
||||
|
||||
# --- Key presence -----------------------------------------------------------
|
||||
|
||||
|
||||
def test_decision_descriptions_has_exactly_10_keys():
|
||||
"""_DECISION_DESCRIPTIONS has exactly 10 keys."""
|
||||
d = _decision_descriptions()
|
||||
assert len(d) == 10
|
||||
|
||||
|
||||
def test_decision_descriptions_keys_match_expected():
|
||||
"""All 10 expected category keys are present, no extras."""
|
||||
d = _decision_descriptions()
|
||||
assert sorted(d.keys()) == _EXPECTED_KEYS
|
||||
|
||||
|
||||
def test_tool_use_agentic_excluded():
|
||||
"""tool_use_agentic must NOT appear in _DECISION_DESCRIPTIONS."""
|
||||
d = _decision_descriptions()
|
||||
assert "tool_use_agentic" not in d
|
||||
|
||||
|
||||
# --- Not a shallow copy of encoder descriptions ------------------------------
|
||||
|
||||
|
||||
def test_decision_descriptions_not_same_object_as_encoder():
|
||||
"""_DECISION_DESCRIPTIONS is a distinct dict, not local_encoder's."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d is not e
|
||||
|
||||
|
||||
def test_decision_descriptions_differ_from_encoder():
|
||||
"""At least one value differs — must not be a shallow copy."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
# They must NOT be equal at the value level
|
||||
assert d != e
|
||||
|
||||
|
||||
def test_coding_general_differ():
|
||||
"""coding_general was tuned — should include 'tracing' wording."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d["coding_general"] != e["coding_general"]
|
||||
assert "tracing" in d["coding_general"]
|
||||
|
||||
|
||||
def test_debugging_differ():
|
||||
"""debugging was tuned — should include 'fixing a bug' wording."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d["debugging"] != e["debugging"]
|
||||
assert "fixing a bug" in d["debugging"]
|
||||
|
||||
|
||||
def test_reasoning_math_differ():
|
||||
"""reasoning_math was tuned — should be a pure word problem, no code."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d["reasoning_math"] != e["reasoning_math"]
|
||||
assert "no code" in d["reasoning_math"]
|
||||
|
||||
|
||||
def test_file_summarization_differ():
|
||||
"""file_summarization was tuned — should emphasize 'single source file'."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d["file_summarization"] != e["file_summarization"]
|
||||
assert "single source file" in d["file_summarization"]
|
||||
|
||||
|
||||
def test_docs_writing_differ():
|
||||
"""docs_writing was tuned — should include '(not code)' qualifier."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d["docs_writing"] != e["docs_writing"]
|
||||
assert "(not code)" in d["docs_writing"]
|
||||
|
||||
|
||||
# --- Keep-same keys still differ from encoder --------------------------------
|
||||
# Even categories marked "keep" differ because the encoder has 11 keys
|
||||
# and decision has 10 — the decision dict is a new dict, so `is` must fail.
|
||||
|
||||
|
||||
def test_unmodified_keys_still_distinct_object():
|
||||
"""summarization key value is the same text but the dict itself differs."""
|
||||
d = _decision_descriptions()
|
||||
e = _encoder_descriptions()
|
||||
assert d is not e
|
||||
|
||||
|
||||
# --- local_encoder.py is NOT modified ----------------------------------------
|
||||
|
||||
|
||||
def test_local_encoder_file_unmodified():
|
||||
"""git diff src/local_encoder.py must be empty — no encoder edits."""
|
||||
repo_root = pathlib.Path(__file__).parent.parent.parent
|
||||
result = subprocess.run(
|
||||
["git", "diff", "src/local_encoder.py"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
cwd=str(repo_root),
|
||||
)
|
||||
assert result.stdout == "", (
|
||||
f"local_encoder.py was modified! Diff:\n{result.stdout}"
|
||||
)
|
||||
|
||||
|
||||
# --- classify_choice compatibility -------------------------------------------
|
||||
|
||||
|
||||
def test_decision_descriptions_keys_can_be_used_as_options():
|
||||
"""The 10 keys are valid option identifiers for classify_choice."""
|
||||
d = _decision_descriptions()
|
||||
# Keys should all be valid strings suitable for option labels
|
||||
for key, value in d.items():
|
||||
assert isinstance(key, str)
|
||||
assert isinstance(value, str)
|
||||
assert len(value) > 0, f"{key!r} has empty description"
|
||||
582
tests/test_local_decision_parse.py
Normal file
582
tests/test_local_decision_parse.py
Normal file
@@ -0,0 +1,582 @@
|
||||
"""Tests for local_decision.parse_logprobs and local_decision.classify_choice.
|
||||
|
||||
Covers (parse_logprobs):
|
||||
- Basic option-letter extraction from Ollama logprobs.
|
||||
- Token stripping: trailing dots, leading whitespace.
|
||||
- Non-option tokens are silently ignored.
|
||||
- Confidence = winner_mass / total_option_mass.
|
||||
- Coverage = total_option_mass (sum of all option letter masses).
|
||||
- RuntimeError on coverage below minimum.
|
||||
- RuntimeError on zero total mass.
|
||||
- RuntimeError on missing logprobs.
|
||||
- Multiple positions and overlapping top_logprobs entries.
|
||||
|
||||
Covers (classify_choice):
|
||||
- POSTs to the Ollama native /api/chat endpoint with the right body.
|
||||
- timeout_s maps to the requests timeout.
|
||||
- Returns (label, confidence, coverage) from parse_logprobs.
|
||||
- Propagates requests.exceptions.RequestException.
|
||||
- Forwards coverage_min to parse_logprobs.
|
||||
"""
|
||||
|
||||
import json
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from local_decision import parse_logprobs
|
||||
|
||||
# --- Fixtures ----------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_logprob(
|
||||
token: str,
|
||||
logprob: float,
|
||||
top_logprobs: list[dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
"""Convenience builder for a logprob position entry."""
|
||||
return {
|
||||
"token": token,
|
||||
"logprob": logprob,
|
||||
"top_logprobs": top_logprobs,
|
||||
}
|
||||
|
||||
|
||||
def _make_response(
|
||||
positions: list[dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
"""Build a minimal Ollama-style response dict from position entries."""
|
||||
return {"logprobs": positions}
|
||||
|
||||
|
||||
# --- Basic parsing -----------------------------------------------------------
|
||||
|
||||
|
||||
def test_simple_two_option_a_wins():
|
||||
"""A has much higher mass than B -> (A, ~1.0, total)."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "B", "logprob": -3.0},
|
||||
],
|
||||
),
|
||||
_make_logprob(
|
||||
".",
|
||||
-0.5,
|
||||
[{"token": ".", "logprob": -0.5}],
|
||||
),
|
||||
])
|
||||
label, confidence, coverage = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "A"
|
||||
assert confidence > 0.95
|
||||
assert coverage > 0
|
||||
|
||||
|
||||
def test_simple_two_option_b_wins():
|
||||
"""B has much higher mass than A -> (B, ~1.0, total)."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"B",
|
||||
-0.02,
|
||||
[
|
||||
{"token": "A", "logprob": -4.0},
|
||||
{"token": "B", "logprob": -0.02},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "B"
|
||||
assert confidence > 0.98
|
||||
|
||||
|
||||
def test_four_options_a_wins():
|
||||
"""A wins among A, B, C, D."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.05,
|
||||
[
|
||||
{"token": "A", "logprob": -0.05},
|
||||
{"token": "B", "logprob": -2.0},
|
||||
{"token": "C", "logprob": -3.0},
|
||||
{"token": "D", "logprob": -4.0},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, confidence, _ = parse_logprobs(response, ["A", "B", "C", "D"])
|
||||
assert label == "A"
|
||||
assert confidence > 0.82
|
||||
|
||||
|
||||
# --- Token stripping ---------------------------------------------------------
|
||||
|
||||
|
||||
def test_trailing_dot_stripped():
|
||||
"""Ollama sends 'A.' as the option separator token — must match 'A'."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A.",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A.", "logprob": -0.01},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, _, _ = parse_logprobs(response, ["A"])
|
||||
assert label == "A"
|
||||
|
||||
|
||||
def test_leading_whitespace_token_ignored():
|
||||
"""Token ' thinking' — after stripping whitespace it is not an option."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
" thinking",
|
||||
-0.1,
|
||||
[
|
||||
{"token": " thinking", "logprob": -0.1},
|
||||
{"token": " response", "logprob": -0.2},
|
||||
],
|
||||
),
|
||||
])
|
||||
# Should not raise; no option letters found so coverage < 0.3.
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_thinking_token_does_not_pollute_mass():
|
||||
"""'thinking' and 'response' tokens are ignored, not counted."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "thinking", "logprob": -0.5},
|
||||
{"token": "response", "logprob": -0.6},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "A"
|
||||
# Only A's mass contributes — confidence should be 1.0
|
||||
assert confidence == pytest.approx(1.0)
|
||||
|
||||
|
||||
# --- Confidence and coverage -------------------------------------------------
|
||||
|
||||
|
||||
def test_confidence_is_normalized_mass_ratio():
|
||||
"""confidence = winner_mass / total_option_mass."""
|
||||
# Use -ln(2) ≈ -0.693 so exp ≈ 0.5 for A, and -ln(1.5) ≈ -0.405 for B
|
||||
logprob_a = math.log(2.0) # ≈ 0.693, exp = 0.5
|
||||
logprob_b = math.log(1.5) # ≈ 0.405, exp = 0.667
|
||||
# A's mass = 0.5, B's mass = 0.667 → total = 1.167, confidence(A) = 0.5/1.167
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"B",
|
||||
-logprob_b,
|
||||
[
|
||||
{"token": "A", "logprob": -logprob_a},
|
||||
{"token": "B", "logprob": -logprob_b},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, confidence, coverage = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "B"
|
||||
expected_confidence = math.exp(-logprob_b) / (math.exp(-logprob_a) + math.exp(-logprob_b))
|
||||
assert confidence == pytest.approx(expected_confidence, rel=1e-9)
|
||||
assert coverage == pytest.approx(math.exp(-logprob_a) + math.exp(-logprob_b), rel=1e-9)
|
||||
|
||||
|
||||
def test_coverage_is_total_mass():
|
||||
"""Coverage = sum of exp(logprob) across all option letters."""
|
||||
logprob_val = -0.1
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
logprob_val,
|
||||
[
|
||||
{"token": "A", "logprob": logprob_val},
|
||||
{"token": "B", "logprob": logprob_val},
|
||||
],
|
||||
),
|
||||
])
|
||||
_, _, coverage = parse_logprobs(response, ["A", "B"])
|
||||
expected = math.exp(logprob_val) + math.exp(logprob_val)
|
||||
assert coverage == pytest.approx(expected, rel=1e-9)
|
||||
|
||||
|
||||
# --- Multiple positions ------------------------------------------------------
|
||||
|
||||
|
||||
def test_multiple_positions_accumulate_mass():
|
||||
"""Logprobs from multiple positions add up per option."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.1,
|
||||
[{"token": "A", "logprob": -0.1}, {"token": "B", "logprob": -0.2}],
|
||||
),
|
||||
_make_logprob(
|
||||
"B",
|
||||
-0.05,
|
||||
[{"token": "A", "logprob": -0.3}, {"token": "B", "logprob": -0.05}],
|
||||
),
|
||||
])
|
||||
label, _, coverage = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "B"
|
||||
# A mass = exp(-0.1) + exp(-0.3), B mass = exp(-0.2) + exp(-0.05)
|
||||
expected_a = math.exp(-0.1) + math.exp(-0.3)
|
||||
expected_b = math.exp(-0.2) + math.exp(-0.05)
|
||||
assert coverage == pytest.approx(expected_a + expected_b, rel=1e-9)
|
||||
|
||||
|
||||
def test_mixed_positions_with_non_option_tokens():
|
||||
"""Mix of option letters and noise tokens across multiple positions."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "thinking", "logprob": -1.0},
|
||||
],
|
||||
),
|
||||
_make_logprob(
|
||||
" response",
|
||||
-0.3,
|
||||
[
|
||||
{"token": "B", "logprob": -0.3},
|
||||
{"token": "assistant", "logprob": -2.0},
|
||||
],
|
||||
),
|
||||
])
|
||||
label, _, _ = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "A" # A: exp(-0.01) ≈ 0.99 > B: exp(-0.3) ≈ 0.74
|
||||
|
||||
|
||||
# --- Coverage threshold errors -----------------------------------------------
|
||||
|
||||
|
||||
def test_runtime_error_on_low_coverage():
|
||||
"""Coverage below minimum raises RuntimeError."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-10.0, # very low probability -> very small mass
|
||||
[{"token": "A", "logprob": -10.0}],
|
||||
),
|
||||
])
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A", "B"], coverage_min=0.3)
|
||||
|
||||
|
||||
def test_runtime_error_on_zero_coverage():
|
||||
"""No option letters in logprobs → zero coverage → raise."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"thinking",
|
||||
-0.1,
|
||||
[
|
||||
{"token": "thinking", "logprob": -0.1},
|
||||
{"token": "response", "logprob": -0.2},
|
||||
],
|
||||
),
|
||||
])
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_runtime_error_on_empty_logprobs():
|
||||
"""Empty logprobs list raises RuntimeError about missing logprobs."""
|
||||
response: dict[str, Any] = {"logprobs": []}
|
||||
with pytest.raises(RuntimeError, match="no logprobs"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_runtime_error_on_missing_logprobs_key():
|
||||
"""Missing logprobs key → default to [] → raise."""
|
||||
response: dict[str, Any] = {}
|
||||
with pytest.raises(RuntimeError, match="no logprobs"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_custom_coverage_min():
|
||||
"""Higher coverage_min requires stronger evidence."""
|
||||
logprob_val = -0.1
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
logprob_val,
|
||||
[{"token": "A", "logprob": logprob_val}],
|
||||
),
|
||||
])
|
||||
# Default min (0.3) should pass since exp(-0.1) ≈ 0.905
|
||||
_, _, _ = parse_logprobs(response, ["A"])
|
||||
# Very high min should fail
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A"], coverage_min=10.0)
|
||||
|
||||
|
||||
# --- Edge cases --------------------------------------------------------------
|
||||
|
||||
|
||||
def test_single_option_letter():
|
||||
"""Single option letter in the list — must still work."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[{"token": "A", "logprob": -0.01}],
|
||||
),
|
||||
])
|
||||
label, confidence, _ = parse_logprobs(response, ["A"])
|
||||
assert label == "A"
|
||||
assert confidence == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_option_letters_not_present_in_response():
|
||||
"""When none of the option letters appear in top_logprobs, coverage is 0."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"X",
|
||||
-0.5,
|
||||
[{"token": "X", "logprob": -0.5}],
|
||||
),
|
||||
])
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_dot_token_only_is_ignored():
|
||||
"""A bare '.' token (not a stripped option) should be ignored."""
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
".",
|
||||
-0.5,
|
||||
[{"token": ".", "logprob": -0.5}],
|
||||
),
|
||||
])
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
parse_logprobs(response, ["A", "B"])
|
||||
|
||||
|
||||
def test_empty_top_logprobs_list():
|
||||
"""Position with an empty top_logprobs list is skipped gracefully."""
|
||||
response = _make_response([
|
||||
_make_logprob("A", -0.1, []),
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[{"token": "A", "logprob": -0.01}],
|
||||
),
|
||||
])
|
||||
label, _, _ = parse_logprobs(response, ["A"])
|
||||
assert label == "A"
|
||||
|
||||
|
||||
def test_missing_top_logprobs_key():
|
||||
"""Position entry missing top_logprobs key does not crash."""
|
||||
response = _make_response([
|
||||
{"token": "A", "logprob": -0.1}, # no top_logprobs key
|
||||
_make_logprob(
|
||||
"B",
|
||||
-0.01,
|
||||
[{"token": "B", "logprob": -0.01}],
|
||||
),
|
||||
])
|
||||
label, _, _ = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "B"
|
||||
|
||||
|
||||
def test_missing_logprob_entry_is_skipped():
|
||||
response = _make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "B"}, # no logprob key
|
||||
],
|
||||
),
|
||||
])
|
||||
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
||||
assert label == "A"
|
||||
# Only A's mass contributes — B's missing logprob is skipped
|
||||
assert confidence == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_fixture_real_response_parses():
|
||||
import pathlib
|
||||
|
||||
fixture_path = pathlib.Path(__file__).parent / "fixtures" / "real_ollama_response.json"
|
||||
with open(fixture_path) as f:
|
||||
response = json.load(f)
|
||||
|
||||
label, confidence, coverage = parse_logprobs(response, ["A", "B", "C", "D"])
|
||||
assert label == "A"
|
||||
assert confidence > 0.9
|
||||
assert coverage > 0
|
||||
|
||||
|
||||
# --- classify_choice --------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload: dict[str, Any]) -> None:
|
||||
self._payload = payload
|
||||
|
||||
def json(self) -> dict[str, Any]:
|
||||
return self._payload
|
||||
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def test_classify_choice_posts_native_api_and_parses(monkeypatch):
|
||||
"""classify_choice POSTs to /api/chat and returns parse_logprobs result."""
|
||||
import local_decision
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def fake_post(url, *, json, timeout):
|
||||
captured["url"] = url
|
||||
captured["json"] = json
|
||||
captured["timeout"] = timeout
|
||||
return _FakeResponse(_make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-0.01,
|
||||
[
|
||||
{"token": "A", "logprob": -0.01},
|
||||
{"token": "B", "logprob": -3.0},
|
||||
],
|
||||
),
|
||||
]))
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
label, confidence, coverage = local_decision.classify_choice(
|
||||
"Write a function",
|
||||
{"A": "writing new code", "B": "refactoring"},
|
||||
base_url="http://ollama:11434",
|
||||
model="qwen3.5:4b",
|
||||
num_ctx=8192,
|
||||
timeout_s=120,
|
||||
)
|
||||
|
||||
assert captured["url"] == "http://ollama:11434/api/chat"
|
||||
assert captured["timeout"] == 120
|
||||
|
||||
payload = captured["json"]
|
||||
assert payload["model"] == "qwen3.5:4b"
|
||||
assert payload["stream"] is False
|
||||
assert payload["think"] is False
|
||||
assert payload["logprobs"] is True
|
||||
assert payload["top_logprobs"] == 20
|
||||
assert payload["options"] == {
|
||||
"num_predict": 1,
|
||||
"temperature": 0,
|
||||
"num_ctx": 8192,
|
||||
}
|
||||
assert payload["messages"][0] == {"role": "system", "content": "answer with the letter only"}
|
||||
user_content = payload["messages"][1]["content"]
|
||||
assert "Write a function" in user_content
|
||||
assert "A. writing new code" in user_content
|
||||
assert "B. refactoring" in user_content
|
||||
assert user_content.endswith("Answer with the letter only.")
|
||||
|
||||
assert label == "A"
|
||||
assert confidence > 0.95
|
||||
assert coverage > 0
|
||||
|
||||
|
||||
def test_classify_choice_sorted_option_letters(monkeypatch):
|
||||
"""option_letters passed to parse_logprobs are sorted."""
|
||||
import local_decision
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def fake_post(url, *, json, timeout):
|
||||
captured["url"] = url
|
||||
captured["json"] = json
|
||||
captured["timeout"] = timeout
|
||||
return _FakeResponse(_make_response([
|
||||
_make_logprob(
|
||||
"B",
|
||||
-0.02,
|
||||
[
|
||||
{"token": "A", "logprob": -4.0},
|
||||
{"token": "B", "logprob": -0.02},
|
||||
],
|
||||
),
|
||||
]))
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
# Unsorted option insertion order; parse must sort to A, B.
|
||||
label, _, _ = local_decision.classify_choice(
|
||||
"Refactor",
|
||||
{"B": "second", "A": "first"},
|
||||
base_url="http://ollama:11434",
|
||||
model="m",
|
||||
num_ctx=4096,
|
||||
timeout_s=30,
|
||||
)
|
||||
assert label == "B"
|
||||
|
||||
|
||||
def test_classify_choice_propagates_request_exception(monkeypatch):
|
||||
"""requests.exceptions.RequestException propagates out of classify_choice."""
|
||||
import requests as real_requests
|
||||
|
||||
import local_decision
|
||||
|
||||
def fake_post(url, *, json, timeout):
|
||||
raise real_requests.exceptions.ConnectionError("boom")
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
with pytest.raises(real_requests.exceptions.RequestException):
|
||||
local_decision.classify_choice(
|
||||
"x",
|
||||
{"A": "a"},
|
||||
base_url="http://ollama:11434",
|
||||
model="m",
|
||||
num_ctx=4096,
|
||||
timeout_s=30,
|
||||
)
|
||||
|
||||
|
||||
def test_classify_choice_coverage_min_passthrough(monkeypatch):
|
||||
"""coverage_min is forwarded to parse_logprobs."""
|
||||
import local_decision
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def fake_post(url, *, json, timeout):
|
||||
captured["timeout"] = timeout
|
||||
return _FakeResponse(_make_response([
|
||||
_make_logprob(
|
||||
"A",
|
||||
-10.0, # very low mass -> fails high coverage_min
|
||||
[{"token": "A", "logprob": -10.0}],
|
||||
),
|
||||
]))
|
||||
|
||||
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
||||
|
||||
with pytest.raises(RuntimeError, match="coverage"):
|
||||
local_decision.classify_choice(
|
||||
"x",
|
||||
{"A": "a"},
|
||||
base_url="http://ollama:11434",
|
||||
model="m",
|
||||
num_ctx=4096,
|
||||
timeout_s=30,
|
||||
coverage_min=10.0,
|
||||
)
|
||||
Reference in New Issue
Block a user