feat(local-encoder): swap from NLI cross-encoding to embedding+centroid classification #95

Merged
alee merged 1 commits from feat/local-encoder-embedding into main 2026-09-19 03:03:37 +00:00
Owner

What

Swap classifier.mode: local_encoder from zero-shot NLI cross-encoding (facebook/bart-large-mnli via transformers.pipeline("zero-shot-classification")) to sentence-embedding + nearest-centroid classification (BAAI/bge-large-en-v1.5 via AutoModel + AutoTokenizer). This is Phase 1 (implementation) of Wave 5.1 in plans/token-waste-waves.md.

Why

The old NLI pipeline cross-encoded the input against every candidate label's hypothesis template — one forward pass per category (9 categories today). That cost multiplies badly for the per-turn delta-detector role planned in later phases. An embedding model does one forward pass for the input, then cheap cosine similarity against precomputed category-description embeddings — cost independent of category count. It also affords a stronger backbone (bge-large-en-v1.5, ~1.3 GB, CPU-viable on the live classifier host).

What changed

src/local_encoder.py — complete backend rewrite

Before After
transformers.pipeline("zero-shot-classification", ...) — cross-encoder, one pass per label AutoTokenizer + AutoModel — one pass for input, cosine similarity against precomputed description embeddings
pipeline returns {"labels": [...], "scores": [...]} directly Mean-pool + L2-normalize token embeddings → dot product with cached description embeddings → softmax with temperature 0.10
Model+device cache: _pipeline_cache Two caches: _model_cache (model+tokenizer) and _desc_cache (precomputed description embeddings), both keyed by (model_id, device)
Confidence: raw NLI score from pipeline Softmax over cosine similarities at T=0.10 — documented rationale in _SOFTMAX_TEMPERATURE

Both public functions (ensure_available, classify_zero_shot) keep their exact signatures. dispatcher.py needs zero changes.

src/config.py — default model + docstring

  • Default model: "facebook/bart-large-mnli" → "BAAI/bge-large-en-v1.5"
  • LocalEncoderConfig docstring updated to describe the embedding+centroid architecture
  • Field comment above model: replaces the gated-deberta narrative with the new rationale and Wave 5.1 reference

Tests

  • tests/test_local_encoder.py: 13 tests (up from 10), all offline — mocks transformers and torch in sys.modules with a minimal MockTensor class that supports the embedding math operations
  • tests/test_classifier_modes_config.py: updated default-model assertion

Scope / not in this PR

This is Phase 1 only. It deliberately does not include:

  • Phase 2 — watching the new classifier against real traffic and re-tuning confidence_threshold (the live config.local.yaml has confidence_threshold: 0.4, which was tuned against the OLD NLI score distribution and is very likely miscalibrated for the new softmax-based scores). That re-tuning is an operator task through the admin portal's Classifier card.
  • Phase 3 — building a delta-detector on top of the embedding path.

Verification

  • src/local_encoder.py — 13 unit tests pass (all offline, mocked)
  • src/config.py — test_config.py: 47 passed
  • Full suite: 2150 passed, 0 failed (95s)
## What Swap `classifier.mode: local_encoder` from zero-shot NLI cross-encoding (`facebook/bart-large-mnli` via `transformers.pipeline("zero-shot-classification")`) to sentence-embedding + nearest-centroid classification (`BAAI/bge-large-en-v1.5` via `AutoModel` + `AutoTokenizer`). This is **Phase 1** (implementation) of Wave 5.1 in `plans/token-waste-waves.md`. ## Why The old NLI pipeline cross-encoded the input against **every** candidate label's hypothesis template — one forward pass per category (9 categories today). That cost multiplies badly for the per-turn delta-detector role planned in later phases. An embedding model does **one** forward pass for the input, then cheap cosine similarity against precomputed category-description embeddings — cost independent of category count. It also affords a stronger backbone (`bge-large-en-v1.5`, ~1.3 GB, CPU-viable on the live classifier host). ## What changed ### `src/local_encoder.py` — complete backend rewrite | Before | After | |---|---| | `transformers.pipeline("zero-shot-classification", ...)` — cross-encoder, one pass per label | `AutoTokenizer` + `AutoModel` — one pass for input, cosine similarity against precomputed description embeddings | | `pipeline` returns `{"labels": [...], "scores": [...]}` directly | Mean-pool + L2-normalize token embeddings → dot product with cached description embeddings → softmax with temperature 0.10 | | Model+device cache: `_pipeline_cache` | Two caches: `_model_cache` (model+tokenizer) and `_desc_cache` (precomputed description embeddings), both keyed by `(model_id, device)` | | Confidence: raw NLI score from pipeline | Softmax over cosine similarities at T=0.10 — documented rationale in `_SOFTMAX_TEMPERATURE` | Both public functions (`ensure_available`, `classify_zero_shot`) keep their exact signatures. `dispatcher.py` needs **zero** changes. ### `src/config.py` — default model + docstring - Default model: `"facebook/bart-large-mnli"` → `"BAAI/bge-large-en-v1.5"` - `LocalEncoderConfig` docstring updated to describe the embedding+centroid architecture - Field comment above `model:` replaces the gated-deberta narrative with the new rationale and Wave 5.1 reference ### Tests - `tests/test_local_encoder.py`: 13 tests (up from 10), all offline — mocks `transformers` and `torch` in `sys.modules` with a minimal `MockTensor` class that supports the embedding math operations - `tests/test_classifier_modes_config.py`: updated default-model assertion ## Scope / not in this PR This is **Phase 1 only**. It deliberately does **not** include: - **Phase 2** — watching the new classifier against real traffic and re-tuning `confidence_threshold` (the live `config.local.yaml` has `confidence_threshold: 0.4`, which was tuned against the OLD NLI score distribution and is **very likely miscalibrated** for the new softmax-based scores). That re-tuning is an operator task through the admin portal's Classifier card. - **Phase 3** — building a delta-detector on top of the embedding path. ## Verification - `src/local_encoder.py` — 13 unit tests pass (all offline, mocked) - `src/config.py` — `test_config.py`: 47 passed - Full suite: **2150 passed, 0 failed** (95s)
alee added 1 commit 2026-09-18 05:51:09 +00:00
Replace transformers.pipeline('zero-shot-classification', bart-large-mnli)
with AutoModel/AutoTokenizer + mean-pool + L2-normalize + cosine
similarity + softmax (BAAI/bge-large-en-v1.5). One forward pass for the
input, cost independent of category count.

Preserves the exact public interface (ensure_available, classify_zero_shot)
— dispatcher.py needs zero changes.

Also updates LocalEncoderConfig's default model and class docstring in
config.py, and rewrites test_local_encoder.py with 13 offline tests
mocking both transformers and torch.

This is Phase 1 of Wave 5.1 in plans/token-waste-waves.md — implementation
only. Phase 2 (re-tuning confidence_threshold against real traffic) and
Phase 3 (delta-detector) are not included.
alee merged commit 74a20a0aed into main 2026-09-19 03:03:37 +00:00
Sign in to join this conversation.
No Reviewers
No Label
1 Participants
Notifications
Due Date
No due date set.
Dependencies

No dependencies set.

Reference: alee/6krrt#95