|
Download README.md from Badalt/intent-classifier-mpnet: direct link, hf CLI and curl.
- Browser
- Download file 5.05 kB
-
https://huggingface.co/Badalt/intent-classifier-mpnet/resolve/main/README.md
- Command line
-
hf download hf://Badalt/intent-classifier-mpnet/README.md
-
curl -L -o README.md https://huggingface.co/Badalt/intent-classifier-mpnet/resolve/main/README.md
5.05 kB
| license: apache-2.0 | |
| base_model: sentence-transformers/paraphrase-multilingual-mpnet-base-v2 | |
| pipeline_tag: text-classification | |
| language: [en, es, fr, de, zh, nl, it, ar] | |
| tags: [intent-classification, logistics, multilingual, open-set, text-classification] | |
| model-index: | |
| - name: intent-classifier-mpnet | |
| results: | |
| - task: {type: text-classification, name: Intent classification} | |
| dataset: {name: synthetic logistics intents (500 rows, private), type: private} | |
| metrics: | |
| - {type: f1, name: macro-F1 (test), value: 0.8848} | |
| - {type: accuracy, name: accuracy (test), value: 0.8961} | |
| # intent-classifier-mpnet | |
| Routes a short user message from a logistics / supply-chain chat assistant to one of **12 intents**, and abstains with **`unknown`** when it is not confident (calibrated reject rule stored in `config.json`). | |
| Fine-tuned from [`sentence-transformers/paraphrase-multilingual-mpnet-base-v2`](https://huggingface.co/sentence-transformers/paraphrase-multilingual-mpnet-base-v2) as a standard `AutoModelForSequenceClassification` (no custom code). | |
| ## Usage | |
| ```python | |
| import torch | |
| from transformers import AutoModelForSequenceClassification, AutoTokenizer | |
| repo = "Badalt/intent-classifier-mpnet" | |
| tok = AutoTokenizer.from_pretrained(repo) | |
| model = AutoModelForSequenceClassification.from_pretrained(repo).eval() | |
| texts = ["whats the eta on LD-55501 pls", "¿Cuántos remolques hay en el patio?"] | |
| with torch.no_grad(): | |
| logits = model(**tok(texts, padding=True, truncation=True, max_length=64, return_tensors="pt")).logits | |
| probs = torch.softmax(logits / model.config.calibration_temperature, dim=-1) | |
| conf, idx = probs.max(-1) | |
| for c, i in zip(conf, idx): | |
| label = "unknown" if c < model.config.unknown_threshold else model.config.id2label[int(i)] | |
| print(label, round(float(c), 3)) | |
| ``` | |
| Reject rule: `unknown` if `max softmax(logits / 0.4622) < 0.5735`. The threshold keeps 95.0% of known-intent validation messages; the temperature was fitted on the same validation set. | |
| ## Labels | |
| `ai_agent_performance`, `appointment_manager`, `chitchat`, `customer_support`, `document_processing`, `knowledge_base`, `orders`, `other`, `shipment_information.analytics`, `shipment_information.disruptions`, `shipment_information.realtime_query`, `yard_management` (+ `unknown` via the reject rule). | |
| ## Training data | |
| 500 labelled synthetic chat messages (12 intents, 30–55 per class; ~13% non-English: ES, FR, DE, ZH, NL, AR, IT; noisy, colloquial). The dataset is confidential and **not** distributed with this model. Split: stratified 70/15/15 (347 / 76 / 77) with near-duplicate paraphrases kept in the same split; seed 42. | |
| ## Training procedure | |
| lr 5e-05, batch 16, ≤15 epochs with early stopping on validation macro-F1 (best epoch 4), linear decay, warmup 0.1, weight decay 0.01, class-weighted cross-entropy, max_length 64, seed 42. Backbone and learning rate chosen by 5-fold cross-validation among 3 multilingual encoders × 2 learning rates. | |
| ## Evaluation | |
| **Held-out test set (n=77, never used for any decision):** macro-F1 **0.885** (95% bootstrap CI 0.791–0.950), accuracy 0.896; English 0.879 vs non-English 1.000 accuracy (n=11). | |
| **5-fold CV on train+val (n=423, out-of-fold):** macro-F1 0.920 (fold mean 0.920 ± 0.028); non-English accuracy 0.942. | |
| | intent | precision | recall | F1 | n | | |
| |---|---|---|---|---| | |
| | `ai_agent_performance` | 1.00 | 1.00 | 1.00 | 6 | | |
| | `appointment_manager` | 0.88 | 1.00 | 0.93 | 7 | | |
| | `chitchat` | 0.83 | 1.00 | 0.91 | 5 | | |
| | `customer_support` | 0.86 | 1.00 | 0.92 | 6 | | |
| | `document_processing` | 1.00 | 0.83 | 0.91 | 6 | | |
| | `knowledge_base` | 0.67 | 1.00 | 0.80 | 6 | | |
| | `orders` | 1.00 | 0.86 | 0.92 | 7 | | |
| | `other` | 1.00 | 0.40 | 0.57 | 5 | | |
| | `shipment_information.analytics` | 1.00 | 0.80 | 0.89 | 10 | | |
| | `shipment_information.disruptions` | 0.83 | 1.00 | 0.91 | 5 | | |
| | `shipment_information.realtime_query` | 0.89 | 1.00 | 0.94 | 8 | | |
| | `yard_management` | 1.00 | 0.83 | 0.91 | 6 | | |
| **Open-set (unknown intents).** A twin model trained without `yard_management` and `document_processing` was tested on known messages plus all messages of those unseen intents: max-softmax reject rule AUROC 0.791, rejection recall 26.2% at 93.8% retention; best scorer (kNN k=1 [mean]) AUROC 0.821. | |
| ## Intended use & limitations | |
| - First-step router for a logistics assistant; low-confidence messages should go to a clarification / fallback path. | |
| - Trained on 500 synthetic rows: expect lower accuracy on real traffic; monitor confidence, unknown-rate and language mix for drift and re-calibrate the threshold on fresh labelled data. | |
| - The three `shipment_information.*` intents are semantically close and account for a large share of errors. | |
| - Non-English evaluation rests on few examples; languages not seen in training (e.g. Arabic, Italian) are untested beyond anecdotes. | |
| - `other` / `chitchat` are trained catch-alls; genuinely new intents are handled only by the confidence threshold. | |
| Training curves and metrics: https://wandb.ai/badalthakur2212-iisc/intent-classifier/runs/ffy90jmz | |