Text classification with transformers
Be able to fine-tune an encoder for classification and evaluate it.
Prerequisites
- DPrecision, recall, F1 and ROCrequired
- EBERT and masked language modellingrequired
Intuition
The recipe:
- Data: texts plus labels, split into train/val/test (stratified). Check the class distribution — 90/10 needs action.
- Model: a pretrained encoder plus a linear head (num_labels).
- Training: lr 2e-5, 3–5 epochs, evaluate F1 on validation every epoch, keep the best (early stopping).
- Evaluation: F1 per class, not just accuracy. The confusion matrix shows which classes get mixed up.
- Error analysis: read 30 misclassified examples. Often the labels are wrong or the boundary is unclear — that is data, not the model.
Imbalance: a weighted loss (class_weight), oversampling the small classes, or adjusting the threshold per class after training.
Code
import numpy as np, torch
from torch.utils.data import DataLoader
from transformers import AutoTokenizer, AutoModelForSequenceClassification
from sklearn.metrics import f1_score, confusion_matrix
name = "KB/bert-base-swedish-cased"
tok = AutoTokenizer.from_pretrained(name)
m = AutoModelForSequenceClassification.from_pretrained(name, num_labels=3)
opt = torch.optim.AdamW(m.parameters(), lr=2e-5)
weights = torch.tensor([1.0, 1.0, 4.0]) # the small class 2 weighs more
loss_fn = torch.nn.CrossEntropyLoss(weight=weights)
def batches(texts, y, bs=16):
for i in range(0, len(texts), bs):
b = tok(texts[i:i+bs], padding=True, truncation=True, max_length=128, return_tensors="pt")
yield b, torch.tensor(y[i:i+bs])
best = 0
for epoch in range(4):
m.train()
for b, yb in batches(tr_x, tr_y):
loss = loss_fn(m(**b).logits, yb); loss.backward(); opt.step(); opt.zero_grad()
m.eval(); pred = []
with torch.no_grad():
for b, _ in batches(va_x, va_y):
pred += m(**b).logits.argmax(-1).tolist()
f1 = f1_score(va_y, pred, average="macro")
print(epoch, round(f1, 3), confusion_matrix(va_y, pred).tolist())
if f1 > best: best = f1; torch.save(m.state_dict(), "best.pt")
Mastery means
- Fine-tunes an encoder for classification with a correct split
- Evaluates with F1 per class and a confusion matrix
- Handles class imbalance
Sign in to do the exercises and build your mastery up.
Sources
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0
- scikit-learn User Guide (BSD-3) — BSD-3-Clause