Project D: a classifier of your own in PyTorch
Be able to train, evaluate and version a classifier of your own on a dataset of your choice.
Prerequisites
- DGit — version controlrequired
- DMNIST from scratchrequired
- DPrecision, recall, F1 and ROCrequired
Intuition
This is a project, not a lesson. You pick a dataset of your own and do the whole chain yourself.
The requirements — all of this has to be in what you hand in:
| # | Requirement |
|---|---|
| 1 | A dataset with at least 3 classes and at least 300 examples |
| 2 | A split into training / validation / test, done before anything else |
| 3 | A baseline (always the most common class, say) |
| 4 | A trained model that beats the baseline |
| 5 | A training curve with the training and validation loss |
| 6 | A confusion matrix and at least three misclassified examples you have looked at |
| 7 | Everything in Git: code, config, random seed, requirements.txt |
| 8 | A README that lets someone else recreate your result |
Requirement 8 is what decides whether the project is finished. Give the repo to a classmate. If they can run it and get the same numbers you are done. If they cannot, something is missing — nearly always the random seed, a dependency version, or a path that only exists on your machine.
Code
# project/train.py — the whole project in one runnable file
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 fix_seed(seed: int):
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
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 "unknown"
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--epochs", type=int, default=10)
ap.add_argument("--lr", type=float, default=1e-3)
ap.add_argument("--out", type=Path, default=Path("runs"))
a = ap.parse_args()
fix_seed(a.seed)
# ... load the data, split it, train, evaluate ...
run = a.out / time.strftime("%Y%m%d-%H%M%S")
run.mkdir(parents=True)
(run / "metadata.json").write_text(json.dumps({
"args": vars(a) | {"out": str(a.out)},
"git_commit": git_commit(),
"python": sys.version.split()[0],
"torch": torch.__version__,
"accuracy_val": val_acc,
"accuracy_test": test_acc,
"baseline": baseline_acc,
"confusion_matrix": confusion_matrix(yte, pred).tolist(),
}, indent=2, ensure_ascii=False), encoding="utf-8")
torch.save(model.state_dict(), run / "model.pt")
print(f"saved in {run} (commit {git_commit()})")
if __name__ == "__main__":
main()
metadata.json is the project's most important file. It ties the result to the commit, the seed and the versions. Six months later it is the only way to know where the 0.913 in your report actually came from.
torch.use_deterministic_algorithms makes the run reproducible on the same machine. On different hardware the last decimal can still differ — which is why you should report a sensible number of decimals, not all of them.
Interactive
The order to work in. Do the steps in turn; do not jump ahead.
| Step | Do | Finished when |
|---|---|---|
| 1 | Pick a dataset and look at 20 examples with your own eyes | you can describe what separates the classes |
| 2 | Split into training/validation/test | the split lives in code with a seed |
| 3 | Run the baseline | you know the number that has to be beaten |
| 4 | Train the simplest model that could work | it beats the baseline |
| 5 | Plot the training curve | you can see whether it over- or underfits |
| 6 | Improve one thing at a time | every change has a measured effect |
| 7 | Error analysis | you can describe what kind of mistakes the model makes |
| 8 | Once on the test set | the number is in metadata.json |
| 9 | Write the README | a classmate can recreate the result |
Four common mistakes in this project:
- Skipping step 1. Nearly every dataset problem is visible if you look at twenty examples.
- Looking at the test set in step 6. It is then used up, and your final number is too optimistic.
- Changing three things at once. Then you do not know which one helped.
- Reporting only the accuracy. Step 7 is what shows that you understood the problem.
Good dataset suggestions: Fashion-MNIST, CIFAR-10, a text dataset from Hugging Face, or — best of all — data you collect yourself about something you care about. The last is harder and considerably more instructive, since you then meet the data quality problems for real.
Mastery means
- Carries out a whole project from data to evaluation
- Versions everything that affects the result
- Reports the error analysis, not just the accuracy
Sign in to do the exercises and build your mastery up.
Sources
- 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