Agentic Intent Router (PyTorch Embedding Centroid Classifier)

A lightweight, zero-shot intent routing layer designed for Autonomous Agentic RAG Systems. Built on top of sentence-transformers/all-MiniLM-L6-v2, this model computes cosine similarity between input query embeddings and precomputed domain prototype centroids to determine whether an agent should dispatch calls to external tools via the Model Context Protocol (MCP) or proceed with standard synthesis.

Key Features

  • Sub-15ms Latency: Operates entirely on CPU or edge GPU inference without requiring expensive LLM forward passes for structural intent detection.
  • Centroid-Based Proximity: Replaces brittle keyword heuristics and broad Natural Language Inference (NLI) prompts with normalized dense vector clustering.
  • MCP Integration: Directly triggers targeted MCP server tools (crm_data, financial_data, or general_knowledge).

Intended Use

This component serves as the upstream routing node in an agentic loop (e.g., LangGraph state machine). Before executing an agent cycle, queries are classified into target operational domains:

Domain Key Target Tool / Server Action
financial_data Financial MCP Server (yfinance) Fetches live market trends and equity tickers
crm_data Enterprise CRM MCP Server (SQLite) Retrieves customer account status and ARR/MRR
general_knowledge Base LLM Context Synthesizes response directly without tool overhead

Technical Architecture

  1. Tokenization & Masking: Query tokens are passed through the 6-layer MiniLM transformer.
  2. Masked Mean Pooling: Outputs are pooled across valid attention masks to create a uniform 384-dimensional representation: $$\mathbf{e} = \frac{\sum (h_i \cdot m_i)}{\sum m_i}$$
  3. L2 Normalization: Embeddings are normalized onto the unit hypersphere: $$\mathbf{\hat{e}} = \frac{\mathbf{e}}{\Vert{}\mathbf{e}\Vert{}_2}$$
  4. Cosine Similarity: The similarity score is derived against normalized prototype cluster centroids: $$\text{Score}(d) = \mathbf{\hat{e}}{\text{query}} \cdot \mathbf{\hat{e}}{\text{centroid}(d)}$$

Usage Example

import torch
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModel

class IntentRouter:
    def __init__(self):
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        self.tokenizer = AutoTokenizer.from_pretrained("sentence-transformers/all-MiniLM-L6-v2")
        self.model = AutoModel.from_pretrained("sentence-transformers/all-MiniLM-L6-v2").to(self.device)
        self.model.eval()

        # Domain prototypes
        self.prototypes = {
            "financial_data": [
                "stock market price, valuation, and ticker performance",
                "equity shares trading and recent market trend"
            ],
            "crm_data": [
                "customer profile and enterprise account status",
                "client subscription details, MRR, and contracts"
            ],
            "general_knowledge": [
                "geography, history, science, and world capitals",
                "general conversational questions and broad topics"
            ]
        }
        self.domain_centroids = {}
        with torch.no_grad():
            for domain, texts in self.prototypes.items():
                embs = self._embed(texts)
                centroid = F.normalize(embs.mean(dim=0, keepdim=True), p=2, dim=1)
                self.domain_centroids[domain] = centroid

    def _embed(self, texts: list[str]) -> torch.Tensor:
        encoded = self.tokenizer(texts, padding=True, truncation=True, return_tensors="pt").to(self.device)
        outputs = self.model(**encoded)
        mask = encoded["attention_mask"].unsqueeze(-1)
        pooled = torch.sum(outputs.last_hidden_state * mask, dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9)
        return F.normalize(pooled, p=2, dim=1)

    def route(self, query: str) -> str:
        with torch.no_grad():
            query_emb = self._embed([query])
            sims = {d: F.cosine_similarity(query_emb, c).item() for d, c in self.domain_centroids.items()}
        return max(sims, key=sims.get)

# Test execution
router = IntentRouter()
print(router.route("What is the current stock price of Apple?")) # Output: financial_data
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support