Hoppa till innehållet
AI-grafen
E· Universitetsprakmodeller· ca 60 min· utvecklande· verifierad 2026-09-20

Textklassificering med transformers

Kunna finjustera en encoder för klassificering och utvärdera den.

Förkunskaper

Intuition

Receptet:

  1. Data: texter + etiketter, delade i train/val/test (stratifierat). Kolla klassfördelningen — 90/10 kräver åtgärd.
  2. Modell: förtränad encoder + linjärt huvud (num_labels).
  3. Träning: lr 2e-5, 3–5 epoker, utvärdera F1 på val varje epok, behåll bästa (tidigt stopp).
  4. Utvärdering: F1 per klass, inte bara accuracy. Förväxlingsmatris visar vilka klasser som blandas.
  5. Felanalys: läs 30 felklassade exempel. Ofta är etiketterna fel eller gränsen otydlig — det är data, inte modellen.

Obalans: viktad loss (class_weight), översampling av små klasser, eller tröskeljustering per klass efter träning.

Kod

import numpy as np, torch
from torch.utils.data import DataLoader
from transformers import AutoTokenizer, AutoModelForSequenceClassification
from sklearn.metrics import f1_score, confusion_matrix

namn = "KB/bert-base-swedish-cased"
tok = AutoTokenizer.from_pretrained(namn)
m = AutoModelForSequenceClassification.from_pretrained(namn, num_labels=3)
opt = torch.optim.AdamW(m.parameters(), lr=2e-5)
vikter = torch.tensor([1.0, 1.0, 4.0])          # liten klass 2 väger tyngre
loss_fn = torch.nn.CrossEntropyLoss(weight=vikter)

def batchar(texter, y, bs=16):
    for i in range(0, len(texter), bs):
        b = tok(texter[i:i+bs], padding=True, truncation=True, max_length=128, return_tensors="pt")
        yield b, torch.tensor(y[i:i+bs])

bast = 0
for epok in range(4):
    m.train()
    for b, yb in batchar(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 batchar(va_x, va_y):
            pred += m(**b).logits.argmax(-1).tolist()
    f1 = f1_score(va_y, pred, average="macro")
    print(epok, round(f1, 3), confusion_matrix(va_y, pred).tolist())
    if f1 > bast: bast = f1; torch.save(m.state_dict(), "bast.pt")

Behärskning innebär

  • Finjusterar en encoder för klassificering med korrekt uppdelning
  • Utvärderar med F1 per klass och förväxlingsmatris
  • Hanterar klassobalans

Logga in för att göra övningarna och bygga upp din behärskning.

Källor

Alla källor och licenser