File size: 5,923 Bytes
2d140d5
 
95cdfc0
2d140d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d0ddd41
2d140d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d0ddd41
 
 
 
 
2d140d5
476b2b3
 
 
 
 
 
 
 
 
2d140d5
 
476b2b3
 
2d140d5
 
476b2b3
2d140d5
476b2b3
2d140d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
476b2b3
2d140d5
 
 
476b2b3
 
 
2d140d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
"""Trains the DefectCNN on the local NEU-DET mirror and saves weights for the Space app.
Reuses the same architecture/hyperparameters as the main notebook's section 3-4
(mit-group8-explainable-defect-detection.ipynb) -- this script exists only
because the notebook itself never persisted a checkpoint to disk."""
import glob
import os
import re
import shutil
from collections import Counter

import numpy as np
import torch
import torch.nn as nn
import torchvision.transforms as T
from PIL import Image
from torch.utils.data import DataLoader, Dataset

from model import CLASS_NAMES, IMAGE_SIZE, MEAN, STD, DefectCNN

SEED = 41   # matches the notebook, so this checkpoint sees the notebook's exact splits
DATA_ROOT = os.path.join("..", "data", "neu-det-src", "IMAGES")
CHECKPOINT_PATH = "model.pt"
EXAMPLES_DIR = "examples"

torch.manual_seed(SEED)
np.random.seed(SEED)

train_transform = T.Compose([
    T.Resize((IMAGE_SIZE, IMAGE_SIZE)),
    T.RandomHorizontalFlip(),
    T.RandomVerticalFlip(),
    T.RandomRotation(8, fill=128),
    T.ToTensor(),
    T.Normalize(MEAN, STD),
])
eval_transform = T.Compose([
    T.Resize((IMAGE_SIZE, IMAGE_SIZE)),
    T.ToTensor(),
    T.Normalize(MEAN, STD),
])


def class_from_filename(path):
    stem = os.path.splitext(os.path.basename(path))[0]
    name = re.sub(r"_\d+$", "", stem)
    if name not in CLASS_NAMES:
        raise ValueError(f"Unrecognized class in filename: {path!r} -> {name!r}")
    return name


class NEUDETSingleLabel(Dataset):
    def __init__(self, paths, labels_by_path, transform):
        self.paths = paths
        self.labels_by_path = labels_by_path
        self.transform = transform

    def __len__(self):
        return len(self.paths)

    def __getitem__(self, index):
        path = self.paths[index]
        image = Image.open(path).convert("RGB")
        label = self.labels_by_path[path]
        target = torch.zeros(len(CLASS_NAMES))
        target[CLASS_NAMES.index(label)] = 1.0
        return self.transform(image), target, [label]


def main():
    image_paths = sorted(glob.glob(os.path.join(DATA_ROOT, "*.jpg")))
    assert image_paths, f"No images found under {DATA_ROOT}"
    labels_by_path = {path: class_from_filename(path) for path in image_paths}
    print(f"Total images found: {len(image_paths)}")
    for name, count in Counter(labels_by_path.values()).items():
        print(f"  {name:16s} {count:4d}")

    # Stratified 70/10/10/10 per class, identical to the notebook's Section 2, so the
    # served checkpoint is trained and scored on exactly the notebook's splits. It is a
    # separate training run, so its weights and test score differ from the notebook's
    # (95.6% here vs 94.1% there). Only train and the final test split are needed here;
    # val is used for checkpoint selection.
    rng = np.random.default_rng(SEED)
    train_paths, val_paths, eval_paths = [], [], []
    for cls in CLASS_NAMES:
        members = sorted(p for p in image_paths if labels_by_path[p] == cls)
        rng.shuffle(members)
        n = len(members)
        a, b, c = int(0.70 * n), int(0.80 * n), int(0.90 * n)
        train_paths += members[:a]
        val_paths   += members[a:b]
        eval_paths  += members[c:]

    train_data = NEUDETSingleLabel(train_paths, labels_by_path, train_transform)
    val_data   = NEUDETSingleLabel(val_paths, labels_by_path, eval_transform)
    eval_data  = NEUDETSingleLabel(eval_paths, labels_by_path, eval_transform)
    train_loader = DataLoader(train_data, batch_size=32, shuffle=True,
                               generator=torch.Generator().manual_seed(SEED))
    val_loader  = DataLoader(val_data, batch_size=32, shuffle=False)
    eval_loader = DataLoader(eval_data, batch_size=32, shuffle=False)
    print(f"Train: {len(train_data)} 路 Val: {len(val_data)} 路 Test: {len(eval_data)}")

    device = "cpu"
    model = DefectCNN(number_of_classes=len(CLASS_NAMES)).to(device)
    criterion = nn.BCEWithLogitsLoss()
    optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-4)
    epochs = 12
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)

    def run_epoch(loader, training):
        model.train(training)
        total_loss, correct, total = 0.0, 0, 0
        for images, targets, _ in loader:
            images, targets = images.to(device), targets.to(device)
            with torch.set_grad_enabled(training):
                logits = model(images)
                loss = criterion(logits, targets)
                if training:
                    optimizer.zero_grad()
                    loss.backward()
                    optimizer.step()
            total_loss += loss.item() * images.size(0)
            preds = logits.argmax(dim=1)
            correct += (preds == targets.argmax(dim=1)).sum().item()
            total += images.size(0)
        return total_loss / total, correct / total

    for epoch in range(1, epochs + 1):
        train_loss, train_acc = run_epoch(train_loader, True)
        scheduler.step()
        val_loss, val_acc = run_epoch(val_loader, False)
        print(f"epoch {epoch:2d}/{epochs} 路 train loss {train_loss:.4f}, acc {train_acc:.1%} "
              f"路 val loss {val_loss:.4f}, acc {val_acc:.1%}")

    test_loss, test_acc = run_epoch(eval_loader, False)
    print(f"final held-out TEST accuracy: {test_acc:.1%}")

    torch.save(model.state_dict(), CHECKPOINT_PATH)
    print(f"Saved checkpoint to {CHECKPOINT_PATH}")

    os.makedirs(EXAMPLES_DIR, exist_ok=True)
    rng2 = np.random.default_rng(SEED + 1)
    for name in CLASS_NAMES:
        candidates = [p for p in eval_paths if labels_by_path[p] == name]
        if candidates:
            chosen = candidates[int(rng2.integers(0, len(candidates)))]
            shutil.copy(chosen, os.path.join(EXAMPLES_DIR, f"{name}.jpg"))
    print(f"Copied example images to {EXAMPLES_DIR}/")


if __name__ == "__main__":
    main()