foursight-intent-e5 / README.md
henhua21's picture
Upload fine-tuned model, calibration and model card
9ebd7f5 verified
|
Raw History Blame Contribute Delete
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.