Theory is necessary but not sufficient. Today we build a complete image classification system the way a professional would โ for example, a model that classifies photos of crop leaves as healthy or diseased. Every step contains decisions that affect whether the final system works in the field.
1. Define the task and collect data#
- Classes: clearly defined and mutually exclusive (or use multi-label if not).
- Coverage: images should reflect deployment conditions โ phones used by real users, lighting at different times of day, backgrounds, disease stages, crop varieties, regions.
- Labels: agreed guidelines; ideally two annotators per image to measure agreement. Label noise directly caps accuracy.
- Quantity: with transfer learning, a few hundred images per class can already give a useful model.
Organise images in the conventional folder layout:
data/
train/healthy/*.jpg train/leaf_blight/*.jpg train/rust/*.jpg
val/... test/...2. Split correctly#
Split by source, not just randomly: if many photos come from the same field or the same plant, keep them in the same split, otherwise the test set contains near-duplicates of training images and accuracy is inflated. Check for exact and near duplicates (perceptual hashing) across splits.
3. Preprocess and augment#
Use the preprocessing expected by the pretrained backbone (resize, crop, normalise with ImageNet mean and standard deviation). Augment only the training set:
import torch
from torchvision import datasets, transforms, models
mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]
train_tf = transforms.Compose([
transforms.RandomResizedCrop(224, scale=(0.6, 1.0)),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(0.3, 0.3, 0.2),
transforms.ToTensor(), transforms.Normalize(mean, std)])
eval_tf = transforms.Compose([
transforms.Resize(256), transforms.CenterCrop(224),
transforms.ToTensor(), transforms.Normalize(mean, std)])
train_ds = datasets.ImageFolder("data/train", train_tf)
val_ds = datasets.ImageFolder("data/val", eval_tf)
train_dl = torch.utils.data.DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4)
val_dl = torch.utils.data.DataLoader(val_ds, batch_size=64, num_workers=4)
print(train_ds.classes, len(train_ds), len(val_ds))4. Choose a backbone and fine-tune#
Start with a strong, efficient pretrained model (ResNet-50, EfficientNet, ConvNeXt-Tiny, or a small ViT; MobileNet for phones). Replace the head, train the head first, then fine-tune the whole network with a lower learning rate.
device = "cuda" if torch.cuda.is_available() else "cpu"
model = models.convnext_tiny(weights=models.ConvNeXt_Tiny_Weights.IMAGENET1K_V1)
model.classifier[2] = torch.nn.Linear(model.classifier[2].in_features, len(train_ds.classes))
model.to(device)
def train(epochs, params, lr):
opt = torch.optim.AdamW(params, lr=lr, weight_decay=0.05)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=epochs * len(train_dl))
for epoch in range(epochs):
model.train()
for x, y in train_dl:
x, y = x.to(device), y.to(device)
loss = torch.nn.functional.cross_entropy(model(x), y, label_smoothing=0.1)
opt.zero_grad(); loss.backward(); opt.step(); sched.step()
print(epoch, "val acc:", evaluate())
@torch.no_grad()
def evaluate():
model.eval(); correct = total = 0
for x, y in val_dl:
pred = model(x.to(device)).argmax(1).cpu()
correct += (pred == y).sum().item(); total += len(y)
return round(correct / total, 4)
for p in model.features.parameters():
p.requires_grad = False
train(3, model.classifier.parameters(), 1e-3) # stage 1: head only
for p in model.parameters():
p.requires_grad = True
train(10, model.parameters(), 1e-4) # stage 2: full fine-tuning5. Evaluate beyond accuracy#
- Confusion matrix โ which diseases are confused?
- Per-class precision and recall โ a rare but dangerous disease needs high recall.
- Calibration โ are confidence scores meaningful? Consider temperature scaling.
- Slices โ performance by phone model, region, lighting, crop variety.
- Test-time augmentation (averaging predictions over flips/crops) can add a little accuracy.
6. Inspect errors#
Look at misclassified images and the highest-confidence mistakes. Typical findings: mislabelled training data, blurry photos, multiple diseases in one image, backgrounds the model latched onto (a shortcut: e.g. all "diseased" photos taken on one field's soil). Tools like Grad-CAM (covered later) show which regions drive a prediction โ useful for detecting shortcuts.
7. Handle "none of the above"#
Deployed models will receive photos that belong to no class (a hand, a car, a blurry image). Options: add an "other/unknown" class with diverse examples, threshold on confidence, or use out-of-distribution detection. Design the user interface to request a retake when confidence is low.
8. Export and monitor#
Export to ONNX, TorchScript or TFLite; quantise for mobile; version the model with its class list and preprocessing; log predictions and confidence in production to detect drift; collect corrected labels from experts to retrain.