Projekt D: egen klassificerare i PyTorch
Kunna träna, utvärdera och versionera en egen klassificerare på ett valfritt dataset.
Förkunskaper
- DGit — versionshanteringkrävs
- DMNIST från grundenkrävs
- DPrecision, recall, F1 och ROCkrävs
Intuition
Det här är ett projekt, inte en lektion. Du väljer ett eget dataset och gör hela kedjan själv.
Kravlista — allt ska finnas i din inlämning:
| # | Krav |
|---|---|
| 1 | Ett dataset med minst 3 klasser och minst 300 exempel |
| 2 | Uppdelning i träning / validering / test, gjord innan något annat |
| 3 | En baslinje (t.ex. alltid vanligaste klassen) |
| 4 | En tränad modell som slår baslinjen |
| 5 | Träningskurva med tränings- och valideringsförlust |
| 6 | Förväxlingsmatris och minst tre felklassade exempel som du tittat på |
| 7 | Allt i Git: kod, konfiguration, slumpfrö, requirements.txt |
| 8 | En README som gör att någon annan kan återskapa ditt resultat |
Krav 8 är det som avgör om projektet är klart. Ge repot till en klasskamrat. Kan hen köra det och få samma siffror är du färdig. Kan hen inte det saknas något — nästan alltid slumpfröet, en beroendeversion eller en sökväg som bara finns på din dator.
Kod
# projekt/train.py — hela projektet i en körbar fil
import argparse, json, random, subprocess, sys, time
from pathlib import Path
import numpy as np, torch, torch.nn as nn
from sklearn.metrics import confusion_matrix, classification_report
def sla_fast_fro(fro: int):
random.seed(fro); np.random.seed(fro); torch.manual_seed(fro)
torch.use_deterministic_algorithms(True, warn_only=True)
def git_commit() -> str:
try:
return subprocess.check_output(["git", "rev-parse", "--short", "HEAD"],
text=True).strip()
except Exception:
return "okänd"
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--fro", type=int, default=0)
ap.add_argument("--epoker", type=int, default=10)
ap.add_argument("--lr", type=float, default=1e-3)
ap.add_argument("--ut", type=Path, default=Path("korningar"))
a = ap.parse_args()
sla_fast_fro(a.fro)
# ... ladda data, dela, träna, utvärdera ...
kor = a.ut / time.strftime("%Y%m%d-%H%M%S")
kor.mkdir(parents=True)
(kor / "metadata.json").write_text(json.dumps({
"argument": vars(a) | {"ut": str(a.ut)},
"git_commit": git_commit(),
"python": sys.version.split()[0],
"torch": torch.__version__,
"traffsakerhet_val": val_acc,
"traffsakerhet_test": test_acc,
"baslinje": baslinje_acc,
"forvaxlingsmatris": confusion_matrix(yte, pred).tolist(),
}, indent=2, ensure_ascii=False), encoding="utf-8")
torch.save(modell.state_dict(), kor / "modell.pt")
print(f"sparat i {kor} (commit {git_commit()})")
if __name__ == "__main__":
main()
metadata.json är projektets viktigaste fil. Den kopplar ihop resultatet med commiten, fröet och versionerna. Sex månader senare är det enda sättet att veta vad siffran 0,913 i din rapport egentligen kom ifrån.
torch.use_deterministic_algorithms gör körningen reproducerbar på samma maskin. På olika hårdvara kan det fortfarande skilja i sista decimalen — därför ska du rapportera med rimligt antal decimaler, inte alla.
Interaktivt
Ordningen att arbeta i. Gör stegen i tur och ordning; hoppa inte framåt.
| Steg | Gör | Klart när |
|---|---|---|
| 1 | Välj dataset och titta på 20 exempel med egna ögon | du kan beskriva vad som skiljer klasserna |
| 2 | Dela i träning/validering/test | uppdelningen ligger i kod med ett frö |
| 3 | Kör baslinjen | du vet vilken siffra som måste slås |
| 4 | Träna den enklaste modell som kan fungera | den slår baslinjen |
| 5 | Rita träningskurvan | du ser om den över- eller underanpassar |
| 6 | Förbättra en sak i taget | varje ändring har en uppmätt effekt |
| 7 | Felanalys | du kan beskriva vilken sorts fel modellen gör |
| 8 | En gång på testmängden | siffran är i metadata.json |
| 9 | Skriv README | en klasskamrat kan återskapa resultatet |
Fyra vanliga fel i det här projektet:
- Hoppa över steg 1. Nästan alla datasetproblem syns om man tittar på tjugo exempel.
- Titta på testmängden i steg 6. Då är den förbrukad, och din slutsiffra är för optimistisk.
- Ändra tre saker samtidigt. Då vet du inte vilken som hjälpte.
- Rapportera bara träffsäkerheten. Steg 7 är det som visar att du förstått problemet.
Bra datasetförslag: Fashion-MNIST, CIFAR-10, ett textdataset från Hugging Face, eller — bäst av allt — egen insamlad data om något du bryr dig om. Det sista är svårare och betydligt mer lärorikt, eftersom du då också möter datakvalitetsproblemen på riktigt.
Behärskning innebär
- Genomför ett helt projekt från data till utvärdering
- Versionerar allt som påverkar resultatet
- Redovisar felanalys, inte bara träffsäkerhet
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- scikit-learn User Guide (BSD-3) — BSD-3-Clause
- Pro Git (Chacon & Straub) — CC BY-NC-SA 3.0