|
Download README.md from henhua21/foursight-intent-e5: direct link, hf CLI and curl.
- Browser
- Download file 5.28 kB
-
https://huggingface.co/henhua21/foursight-intent-e5/resolve/main/README.md
- Command line
-
hf download hf://henhua21/foursight-intent-e5/README.md
-
curl -L -o README.md https://huggingface.co/henhua21/foursight-intent-e5/resolve/main/README.md
5.28 kB
| license: mit | |
| base_model: intfloat/multilingual-e5-base | |
| pipeline_tag: text-classification | |
| language: [en, es, pt, fr, it, nl, pl, zh, ja] | |
| tags: [intent-classification, multi-turn, logistics] | |
| # FourSight intent classifier (multilingual-e5, with conversation context) | |
| Routes a question asked to the FourSight logistics chat assistant to one of seven intents: | |
| | id | intent | meaning | | |
| |---|---|---| | |
| | 0 | shipment_information | where a shipment is right now (visibility) | | |
| | 1 | realtime_query | live lookups: ETAs, statuses, reference numbers of specific loads | | |
| | 2 | analytics | aggregated answers over historic data (counts, trends) | | |
| | 3 | disruptions | weather, port, border or capacity disruptions | | |
| | 4 | chitchat | small talk with the bot | | |
| | 5 | knowledge_base | how a product or feature works | | |
| | 6 | other | out of scope | | |
| Code, training and evaluation: https://github.com/genzwildsoul0824/foursight-intent-classifier | |
| Training runs: https://wandb.ai/henhua21-tiktok/foursight-intent | |
| ## How to use | |
| The model reads the current question together with the previous two questions of the | |
| conversation, as a text pair. `calibration.json` holds the post-processing that turns the | |
| logits into the final label and confidence: a per-class bias, a transition prior on the | |
| previous turn's predicted intent, and a temperature. | |
| The simplest way is the repo's `predict.py`, which does all of this: | |
| ```bash | |
| git clone https://github.com/genzwildsoul0824/foursight-intent-classifier | |
| cd foursight-intent-classifier && pip install -r requirements.txt | |
| # put test.csv in data/ | |
| cd src && python predict.py --model henhua21/foursight-intent-e5 | |
| ``` | |
| Raw model only (no calibration): | |
| ```python | |
| from transformers import AutoModelForSequenceClassification, AutoTokenizer | |
| import torch | |
| name = "henhua21/foursight-intent-e5" | |
| tok = AutoTokenizer.from_pretrained(name) | |
| model = AutoModelForSequenceClassification.from_pretrained(name).eval() | |
| question = "what about last month?" | |
| context = "how many loads were delivered late in March" # previous turn(s), most recent first, joined with " | " | |
| with torch.no_grad(): | |
| logits = model(**tok(question, context, truncation="longest_first", max_length=128, return_tensors="pt")).logits | |
| print(model.config.id2label[int(logits.argmax())]) | |
| ``` | |
| Without the calibration, the raw model over-predicts the rare classes (see below), so use | |
| `predict.py` for real decisions. | |
| ## Training data | |
| 1,904 anonymised user questions from the FourSight assistant, each labelled with one intent. | |
| The data is not public. It is imbalanced: realtime_query 782, knowledge_base 509, analytics | |
| 257, shipment_information 256, other 55, disruptions 28, chitchat 17. It is multilingual, | |
| mostly English. Rows are in conversation order, which is where the context comes from. | |
| ## Training | |
| - Base model `intfloat/multilingual-e5-base`, with a new 7-class head. Word embeddings are frozen. | |
| - Input is (question, previous two questions) as a text pair, max 128 tokens. Context dropout | |
| 0.5: during training, half the time the model sees the question alone. | |
| - 8 epochs, batch 16, AdamW lr 3e-5, 10% warmup then linear decay, weight decay 0.01. | |
| - Cross-entropy with balanced class weights and label smoothing 0.1. bf16 autocast, fp32 weights. | |
| - Seed 42 with deterministic CUDA algorithms. This model was trained on all 1,904 rows. The | |
| numbers below come from 5-fold cross-validation of the same recipe. | |
| ## Evaluation | |
| 5-fold cross-validation, grouped so that duplicate questions and neighbouring turns of a | |
| conversation never fall in both the training and the validation folds. The post-processing | |
| is tuned inside the cross-validation, so the scores are not optimistic. The score is the | |
| challenge metric, OVERALL = 0.6 macro-F1 + 0.4 accuracy. | |
| | | OVERALL | macro-F1 | accuracy | ECE | | |
| |---|---|---|---|---| | |
| | raw model | 59.3 | 52.9 | 68.8 | 0.210 | | |
| | with `calibration.json` | **67.7** | **63.8** | **73.5** | **0.035** | | |
| | class | precision | recall | F1 | | |
| |---|---|---|---| | |
| | shipment_information | 44.6 | 30.5 | 36.2 | | |
| | realtime_query | 74.1 | 83.5 | 78.5 | | |
| | analytics | 68.8 | 57.6 | 62.7 | | |
| | disruptions | 74.2 | 82.1 | 78.0 | | |
| | chitchat | 75.0 | 88.2 | 81.1 | | |
| | knowledge_base | 86.7 | 93.1 | 89.8 | | |
| | other | 25.7 | 16.4 | 20.0 | | |
| **needs_review:** flagging predictions with calibrated confidence below 0.5 flags 13% of | |
| questions. 60% of the flagged ones are wrong (against 26.5% overall), the flag catches 29% of | |
| all errors, and accuracy on the unflagged 87% is 78.5%. | |
| ## Intended use and limitations | |
| - Intended as the first routing step of the FourSight assistant, with low-confidence | |
| questions sent to a fallback or to human review. Not intended for other domains. | |
| - shipment_information vs realtime_query is the weakest boundary. The training labels are | |
| inconsistent there too: the same load ID appears under both labels. | |
| - "other" is a small catch-all class (55 examples) and is often missed. | |
| - Context comes from the previous rows. Session boundaries were not available in the training | |
| data, so a conversation's first question sometimes gets context from another conversation. | |
| With real session IDs this goes away. | |
| - Deterministic: the same input gives the same label (argmax, no sampling). CPU and GPU gave | |
| identical labels on the test set. | |