Skip to content
AI-grafen
EUniversityLanguage models· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

Text classification with transformers

Be able to fine-tune an encoder for classification and evaluate it.

Prerequisites

Intuition

The recipe:

  1. Data: texts plus labels, split into train/val/test (stratified). Check the class distribution — 90/10 needs action.
  2. Model: a pretrained encoder plus a linear head (num_labels).
  3. Training: lr 2e-5, 3–5 epochs, evaluate F1 on validation every epoch, keep the best (early stopping).
  4. Evaluation: F1 per class, not just accuracy. The confusion matrix shows which classes get mixed up.
  5. 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

All the sources and licences