Dataloaders, batchar och epoker
Kunna bygga Dataset och DataLoader och förstå batchstorlekens effekt.
Förkunskaper
Intuition
Epok = en genomgång av hela träningsdatan. Batch = de exempel som behandlas tillsammans innan vikterna uppdateras. Iteration = en batch.
1 000 exempel med batchstorlek 32 → 32 iterationer per epok (sista batchen har 8).
Batchstorlekens effekt:
| Liten batch (8–32) | Stor batch (256+) |
|---|---|
| brusiga gradienter (kan hjälpa generalisering) | stabila gradienter |
| mindre minne | mer minne — ofta den hårda gränsen |
| fler uppdateringar per epok | snabbare per epok på GPU |
Shuffle varje epok är viktigt: ligger datan sorterad efter klass ser modellen först bara klass 0, sedan bara klass 1, och gradienterna drar åt olika håll i varje halva.
Kod
import torch
from torch.utils.data import Dataset, DataLoader
class TabellData(Dataset):
def __init__(self, X, y):
self.X = torch.tensor(X, dtype=torch.float32)
self.y = torch.tensor(y, dtype=torch.long)
def __len__(self):
return len(self.y)
def __getitem__(self, i):
return self.X[i], self.y[i]
train = DataLoader(TabellData(X_tr, y_tr), batch_size=32, shuffle=True, num_workers=2, drop_last=False)
val = DataLoader(TabellData(X_va, y_va), batch_size=256, shuffle=False)
for epok in range(3):
modell.train()
for xb, yb in train: # xb: (32, features)
loss = kriterium(modell(xb), yb)
loss.backward(); opt.step(); opt.zero_grad()
modell.eval()
with torch.no_grad():
acc = sum((modell(xb).argmax(1) == yb).sum().item() for xb, yb in val) / len(val.dataset)
print(epok, round(acc, 3))
Tre vanliga fel: shuffle=True på valideringsdata (onödigt, försvårar jämförelse), glömt model.eval() och torch.no_grad() vid utvärdering (dropout aktiv och onödigt minne), och för stor batch som ger «CUDA out of memory» — lös med gradient accumulation i stället för att sänka lr.
Behärskning innebär
- Bygger Dataset och DataLoader
- Förklarar batchstorlekens effekt på minne och träning
- Vet vad shuffle och epok betyder
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- PyTorch — Data loading (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0