360 lines
13 KiB
Python
360 lines
13 KiB
Python
"""
|
|
Script 2 — Classification Model Trainer
|
|
=========================================
|
|
Trains a lightweight EfficientNet-B0 (ImageNet pre-trained) on the
|
|
crop dataset produced by script 1.
|
|
|
|
Expected dataset layout
|
|
-----------------------
|
|
crops_dataset/
|
|
train/
|
|
cola_can/
|
|
pepsi_can/
|
|
lays_chips/
|
|
…
|
|
val/
|
|
cola_can/
|
|
…
|
|
test/ (optional)
|
|
…
|
|
|
|
Outputs
|
|
-------
|
|
runs/classify/
|
|
best.pt ← best checkpoint (val accuracy)
|
|
last.pt ← last checkpoint
|
|
class_names.json ← {idx: class_name} mapping used at inference
|
|
training_curves.png
|
|
|
|
Usage
|
|
-----
|
|
python 2_train_classifier.py \
|
|
--data_dir crops_dataset \
|
|
--output_dir runs/classify \
|
|
--model efficientnet_b0 \
|
|
--epochs 50 \
|
|
--batch_size 64 \
|
|
--img_size 224 \
|
|
--lr 1e-3 \
|
|
--patience 10 \
|
|
--workers 4 \
|
|
--amp
|
|
"""
|
|
|
|
import argparse
|
|
import json
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
from torch.utils.data import DataLoader
|
|
from torchvision import datasets, models, transforms
|
|
from torchvision.models import (
|
|
EfficientNet_B0_Weights,
|
|
EfficientNet_B2_Weights,
|
|
MobileNet_V3_Small_Weights,
|
|
ResNet50_1x64d_QuantizedWeights,
|
|
)
|
|
from torch.cuda.amp import GradScaler, autocast
|
|
import matplotlib.pyplot as plt
|
|
from tqdm import tqdm
|
|
|
|
|
|
# ─────────────────────────── args ────────────────────────────────────────────
|
|
|
|
SUPPORTED_MODELS = [
|
|
"efficientnet_b0",
|
|
"efficientnet_b2",
|
|
"mobilenet_v3_small",
|
|
"resnet50",
|
|
]
|
|
|
|
|
|
def parse_args():
|
|
p = argparse.ArgumentParser(description="Train a product classifier on crop dataset")
|
|
p.add_argument("--data_dir", required=True, help="Root of crops_dataset/ with train/ val/ subfolders")
|
|
p.add_argument("--output_dir", default="runs/classify", help="Where to save checkpoints & logs")
|
|
p.add_argument("--model", default="efficientnet_b0", choices=SUPPORTED_MODELS)
|
|
p.add_argument("--epochs", type=int, default=50)
|
|
p.add_argument("--batch_size", type=int, default=64)
|
|
p.add_argument("--img_size", type=int, default=224)
|
|
p.add_argument("--lr", type=float, default=1e-3)
|
|
p.add_argument("--weight_decay",type=float, default=1e-4)
|
|
p.add_argument("--patience", type=int, default=10, help="Early-stop patience (epochs)")
|
|
p.add_argument("--workers", type=int, default=4)
|
|
p.add_argument("--amp", action="store_true", help="Use Automatic Mixed Precision (CUDA only)")
|
|
p.add_argument("--freeze_backbone", action="store_true",
|
|
help="Freeze backbone, only train the classifier head")
|
|
p.add_argument("--seed", type=int, default=42)
|
|
return p.parse_args()
|
|
|
|
|
|
# ─────────────────────────── model factory ───────────────────────────────────
|
|
|
|
def build_model(name: str, num_classes: int, freeze_backbone: bool) -> nn.Module:
|
|
if name == "efficientnet_b0":
|
|
m = models.efficientnet_b0(weights=EfficientNet_B0_Weights.IMAGENET1K_V1)
|
|
if freeze_backbone:
|
|
for p in m.features.parameters():
|
|
p.requires_grad = False
|
|
in_features = m.classifier[1].in_features
|
|
m.classifier[1] = nn.Linear(in_features, num_classes)
|
|
|
|
elif name == "efficientnet_b2":
|
|
m = models.efficientnet_b2(weights=EfficientNet_B2_Weights.IMAGENET1K_V1)
|
|
if freeze_backbone:
|
|
for p in m.features.parameters():
|
|
p.requires_grad = False
|
|
in_features = m.classifier[1].in_features
|
|
m.classifier[1] = nn.Linear(in_features, num_classes)
|
|
|
|
elif name == "mobilenet_v3_small":
|
|
m = models.mobilenet_v3_small(weights=MobileNet_V3_Small_Weights.IMAGENET1K_V1)
|
|
if freeze_backbone:
|
|
for p in m.features.parameters():
|
|
p.requires_grad = False
|
|
in_features = m.classifier[3].in_features
|
|
m.classifier[3] = nn.Linear(in_features, num_classes)
|
|
|
|
elif name == "resnet50":
|
|
m = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
|
|
if freeze_backbone:
|
|
for name_, p in m.named_parameters():
|
|
if "fc" not in name_:
|
|
p.requires_grad = False
|
|
m.fc = nn.Linear(m.fc.in_features, num_classes)
|
|
|
|
else:
|
|
raise ValueError(f"Unknown model: {name}")
|
|
|
|
return m
|
|
|
|
|
|
# ─────────────────────────── data ────────────────────────────────────────────
|
|
|
|
def build_loaders(data_dir: Path, img_size: int, batch_size: int, workers: int):
|
|
mean = [0.485, 0.456, 0.406]
|
|
std = [0.229, 0.224, 0.225]
|
|
|
|
train_tf = transforms.Compose([
|
|
transforms.RandomResizedCrop(img_size, scale=(0.7, 1.0)),
|
|
transforms.RandomHorizontalFlip(),
|
|
transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.05),
|
|
transforms.RandomGrayscale(p=0.05),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean, std),
|
|
])
|
|
val_tf = transforms.Compose([
|
|
transforms.Resize(int(img_size * 1.15)),
|
|
transforms.CenterCrop(img_size),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean, std),
|
|
])
|
|
|
|
train_set = datasets.ImageFolder(str(data_dir / "train"), transform=train_tf)
|
|
val_set = datasets.ImageFolder(str(data_dir / "val"), transform=val_tf)
|
|
|
|
# ── Weighted sampler — balances classes regardless of image count per class.
|
|
# Rare classes get oversampled, dominant ones get undersampled.
|
|
# If your dataset is already balanced this is a no-op in effect.
|
|
import numpy as np
|
|
from torch.utils.data import WeightedRandomSampler
|
|
targets = train_set.targets
|
|
class_counts = np.bincount(targets)
|
|
class_weights = 1.0 / class_counts
|
|
sample_weights = class_weights[targets]
|
|
sampler = WeightedRandomSampler(
|
|
weights = sample_weights,
|
|
num_samples = len(train_set),
|
|
replacement = True,
|
|
)
|
|
|
|
# ── Imbalance report ──────────────────────────────────────────────────────
|
|
ratio = class_counts.max() / max(class_counts.min(), 1)
|
|
if ratio > 3:
|
|
print(f" ⚠ Class imbalance detected — {ratio:.1f}x ratio "
|
|
f"(max {class_counts.max()} / min {class_counts.min()} images).")
|
|
print(" WeightedRandomSampler enabled to compensate.")
|
|
low = [train_set.classes[i] for i, c in enumerate(class_counts) if c < 20]
|
|
if low:
|
|
print(f" ⚠ Classes with <20 training images: {low}")
|
|
print(" Consider collecting more data or using --freeze_backbone.")
|
|
|
|
train_loader = DataLoader(train_set, batch_size=batch_size, sampler=sampler,
|
|
num_workers=workers, pin_memory=True)
|
|
val_loader = DataLoader(val_set, batch_size=batch_size, shuffle=False,
|
|
num_workers=workers, pin_memory=True)
|
|
|
|
return train_loader, val_loader, train_set.classes
|
|
|
|
|
|
# ─────────────────────────── training loop ───────────────────────────────────
|
|
|
|
def run_epoch(model, loader, criterion, optimizer, device, scaler, train: bool):
|
|
model.train(train)
|
|
total_loss, correct, total = 0.0, 0, 0
|
|
|
|
ctx = torch.enable_grad() if train else torch.no_grad()
|
|
with ctx:
|
|
for imgs, labels in tqdm(loader, desc="train" if train else "val ", leave=False):
|
|
imgs, labels = imgs.to(device), labels.to(device)
|
|
|
|
with autocast(enabled=(scaler is not None)):
|
|
logits = model(imgs)
|
|
loss = criterion(logits, labels)
|
|
|
|
if train:
|
|
optimizer.zero_grad(set_to_none=True)
|
|
if scaler:
|
|
scaler.scale(loss).backward()
|
|
scaler.step(optimizer)
|
|
scaler.update()
|
|
else:
|
|
loss.backward()
|
|
optimizer.step()
|
|
|
|
total_loss += loss.item() * imgs.size(0)
|
|
correct += (logits.argmax(1) == labels).sum().item()
|
|
total += imgs.size(0)
|
|
|
|
return total_loss / total, correct / total
|
|
|
|
|
|
def plot_curves(history: dict, out_path: Path):
|
|
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
|
|
epochs = range(1, len(history["train_loss"]) + 1)
|
|
|
|
ax1.plot(epochs, history["train_loss"], label="Train")
|
|
ax1.plot(epochs, history["val_loss"], label="Val")
|
|
ax1.set_title("Loss"); ax1.set_xlabel("Epoch"); ax1.legend()
|
|
|
|
ax2.plot(epochs, history["train_acc"], label="Train")
|
|
ax2.plot(epochs, history["val_acc"], label="Val")
|
|
ax2.set_title("Accuracy"); ax2.set_xlabel("Epoch"); ax2.legend()
|
|
|
|
fig.tight_layout()
|
|
fig.savefig(str(out_path), dpi=150)
|
|
plt.close(fig)
|
|
print(f" Curves saved → {out_path}")
|
|
|
|
|
|
# ─────────────────────────── main ────────────────────────────────────────────
|
|
|
|
def main():
|
|
args = parse_args()
|
|
torch.manual_seed(args.seed)
|
|
|
|
out_dir = Path(args.output_dir)
|
|
out_dir.mkdir(parents=True, exist_ok=True)
|
|
data_dir = Path(args.data_dir)
|
|
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
print(f"Device: {device}")
|
|
|
|
# Data
|
|
print("Loading dataset …")
|
|
train_loader, val_loader, class_names = build_loaders(
|
|
data_dir, args.img_size, args.batch_size, args.workers
|
|
)
|
|
num_classes = len(class_names)
|
|
print(f" Classes ({num_classes}): {class_names}")
|
|
|
|
# Save class names
|
|
class_map = {i: c for i, c in enumerate(class_names)}
|
|
with open(out_dir / "class_names.json", "w") as f:
|
|
json.dump(class_map, f, indent=2)
|
|
print(f" class_names.json → {out_dir / 'class_names.json'}")
|
|
|
|
# Model
|
|
print(f"\nBuilding model: {args.model}")
|
|
model = build_model(args.model, num_classes, args.freeze_backbone).to(device)
|
|
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
|
total = sum(p.numel() for p in model.parameters())
|
|
print(f" Trainable params: {trainable:,} / {total:,}")
|
|
|
|
# Optimiser & scheduler
|
|
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
|
|
optimizer = torch.optim.AdamW(
|
|
filter(lambda p: p.requires_grad, model.parameters()),
|
|
lr=args.lr, weight_decay=args.weight_decay
|
|
)
|
|
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
|
|
scaler = GradScaler() if (args.amp and device.type == "cuda") else None
|
|
|
|
# Training
|
|
history = {"train_loss": [], "train_acc": [], "val_loss": [], "val_acc": []}
|
|
best_val_acc = 0.0
|
|
patience_ctr = 0
|
|
|
|
print(f"\n{'─'*55}")
|
|
print(f" Starting training for {args.epochs} epochs")
|
|
print(f"{'─'*55}")
|
|
|
|
for epoch in range(1, args.epochs + 1):
|
|
t0 = time.time()
|
|
|
|
train_loss, train_acc = run_epoch(model, train_loader, criterion, optimizer, device, scaler, train=True)
|
|
val_loss, val_acc = run_epoch(model, val_loader, criterion, None, device, None, train=False)
|
|
|
|
scheduler.step()
|
|
elapsed = time.time() - t0
|
|
|
|
history["train_loss"].append(train_loss)
|
|
history["train_acc"].append(train_acc)
|
|
history["val_loss"].append(val_loss)
|
|
history["val_acc"].append(val_acc)
|
|
|
|
improved = val_acc > best_val_acc
|
|
tag = " ✓ best" if improved else ""
|
|
print(
|
|
f"Epoch {epoch:3d}/{args.epochs} | "
|
|
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
|
|
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
|
|
f"{elapsed:.1f}s{tag}"
|
|
)
|
|
|
|
# Save checkpoints
|
|
if improved:
|
|
best_val_acc = val_acc
|
|
patience_ctr = 0
|
|
torch.save({
|
|
"epoch": epoch,
|
|
"model_state": model.state_dict(),
|
|
"class_names": class_names,
|
|
"val_acc": val_acc,
|
|
"args": vars(args),
|
|
}, out_dir / "best.pt")
|
|
else:
|
|
patience_ctr += 1
|
|
|
|
torch.save({
|
|
"epoch": epoch,
|
|
"model_state": model.state_dict(),
|
|
"class_names": class_names,
|
|
"val_acc": val_acc,
|
|
"args": vars(args),
|
|
}, out_dir / "last.pt")
|
|
|
|
# Early stopping
|
|
if patience_ctr >= args.patience:
|
|
print(f"\n Early stopping triggered (no improvement for {args.patience} epochs)")
|
|
break
|
|
|
|
plot_curves(history, out_dir / "training_curves.png")
|
|
|
|
print("\n" + "═" * 55)
|
|
print(" Training complete!")
|
|
print(f" Best val accuracy : {best_val_acc:.4f}")
|
|
print(f" Checkpoints saved : {out_dir}")
|
|
print("═" * 55)
|
|
print("\n ➜ Next step:")
|
|
print(" python 3_inference.py \\")
|
|
print(f" --weights {out_dir / 'best.pt'} \\")
|
|
print(" --source /path/to/images \\")
|
|
print(" --detector_weights /path/to/best_yolo.pt")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|