feat(local-encoder): swap from NLI cross-encoding to embedding+centroid classification #95
Reference in New Issue
Block a user
Delete Branch "feat/local-encoder-embedding"
Deleting a branch is permanent. Although the deleted branch may continue to exist for a short time before it actually gets removed, it CANNOT be undone in most cases. Continue?
What
Swap
classifier.mode: local_encoderfrom zero-shot NLI cross-encoding (facebook/bart-large-mnliviatransformers.pipeline("zero-shot-classification")) to sentence-embedding + nearest-centroid classification (BAAI/bge-large-en-v1.5viaAutoModel+AutoTokenizer). This is Phase 1 (implementation) of Wave 5.1 inplans/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 rewritetransformers.pipeline("zero-shot-classification", ...)— cross-encoder, one pass per labelAutoTokenizer+AutoModel— one pass for input, cosine similarity against precomputed description embeddingspipelinereturns{"labels": [...], "scores": [...]}directly_pipeline_cache_model_cache(model+tokenizer) and_desc_cache(precomputed description embeddings), both keyed by(model_id, device)_SOFTMAX_TEMPERATUREBoth public functions (
ensure_available,classify_zero_shot) keep their exact signatures.dispatcher.pyneeds zero changes.src/config.py— default model + docstring"facebook/bart-large-mnli"→"BAAI/bge-large-en-v1.5"LocalEncoderConfigdocstring updated to describe the embedding+centroid architecturemodel:replaces the gated-deberta narrative with the new rationale and Wave 5.1 referenceTests
tests/test_local_encoder.py: 13 tests (up from 10), all offline — mockstransformersandtorchinsys.moduleswith a minimalMockTensorclass that supports the embedding math operationstests/test_classifier_modes_config.py: updated default-model assertionScope / not in this PR
This is Phase 1 only. It deliberately does not include:
confidence_threshold(the liveconfig.local.yamlhasconfidence_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.Verification
src/local_encoder.py— 13 unit tests pass (all offline, mocked)src/config.py—test_config.py: 47 passedReplace 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 referenced this pull request2026-09-19 04:24:41 +00:00