Skip to content
AI-grafen
DAI developerDeep learning· about 45 min· fundamentals that rarely change· verified 2026-09-20· EN

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

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
1A dataset with at least 3 classes and at least 300 examples
2A split into training / validation / test, done before anything else
3A baseline (always the most common class, say)
4A trained model that beats the baseline
5A training curve with the training and validation loss
6A confusion matrix and at least three misclassified examples you have looked at
7Everything in Git: code, config, random seed, requirements.txt
8A 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.

StepDoFinished when
1Pick a dataset and look at 20 examples with your own eyesyou can describe what separates the classes
2Split into training/validation/testthe split lives in code with a seed
3Run the baselineyou know the number that has to be beaten
4Train the simplest model that could workit beats the baseline
5Plot the training curveyou can see whether it over- or underfits
6Improve one thing at a timeevery change has a measured effect
7Error analysisyou can describe what kind of mistakes the model makes
8Once on the test setthe number is in metadata.json
9Write the READMEa classmate can recreate the result

Four common mistakes in this project:

  1. Skipping step 1. Nearly every dataset problem is visible if you look at twenty examples.
  2. Looking at the test set in step 6. It is then used up, and your final number is too optimistic.
  3. Changing three things at once. Then you do not know which one helped.
  4. 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

All the sources and licences