--- 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.