-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
259 lines (217 loc) · 11.1 KB
/
Copy pathtrain.py
File metadata and controls
259 lines (217 loc) · 11.1 KB
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
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
"""Train the network on MNIST.
python train.py # sensible defaults, ~98.3% test
python train.py --epochs 40 --hidden 512 256
python train.py --optimizer sgd --lr 0.1 --no-batchnorm
Model selection is done on the validation split; the test set is touched once,
at the very end, using the best-validation checkpoint. Peeking at test between
epochs and keeping the best one is the classic way to report a number you
cannot reproduce.
"""
import argparse
import json
import os
import time
import numpy as np
import augment as augmentation
import data
from nn import (Adam, CosineSchedule, SGD, SoftmaxCrossEntropy, build_model,
checkpoint_path, load_checkpoint, save_checkpoint)
HERE = os.path.dirname(os.path.abspath(__file__))
def parse_args():
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--arch", choices=["mlp", "cnn"], default="mlp")
p.add_argument("--epochs", type=int, default=None,
help="default 20 for mlp, 8 for cnn")
p.add_argument("--batch-size", type=int, default=128)
p.add_argument("--lr", type=float, default=1e-3)
p.add_argument("--optimizer", choices=["adam", "sgd"], default="adam")
p.add_argument("--momentum", type=float, default=0.9, help="sgd only")
p.add_argument("--weight-decay", type=float, default=0.0)
p.add_argument("--hidden", type=int, nargs="+", default=[256, 128],
help="mlp: hidden widths")
p.add_argument("--channels", type=int, nargs="+", default=[16, 32],
help="cnn: channels per conv block")
p.add_argument("--kernel-size", type=int, default=5, help="cnn")
p.add_argument("--fc-hidden", type=int, nargs="+", default=[128],
help="cnn: dense head widths")
p.add_argument("--activation", choices=["relu", "tanh"], default="relu")
p.add_argument("--dropout", type=float, default=0.2)
p.add_argument("--no-batchnorm", dest="batchnorm", action="store_false")
p.add_argument("--augment", action="store_true",
help="random affine augmentation (train split only)")
p.add_argument("--aug-rotation", type=float, default=10.0, help="degrees")
p.add_argument("--aug-translate", type=float, default=2.0, help="pixels")
p.add_argument("--aug-scale", type=float, default=0.1, help="fraction")
p.add_argument("--elastic", action="store_true",
help="elastic distortion; combines with --augment")
p.add_argument("--elastic-alpha", type=float, default=1.5,
help="RMS displacement in pixels")
p.add_argument("--elastic-sigma", type=float, default=6.0,
help="displacement-field blur, pixels")
p.add_argument("--warmup", type=int, default=100, help="warmup steps")
p.add_argument("--val-size", type=int, default=5000)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--out", default=None, help="default checkpoints/<arch>.npz")
p.add_argument("--plot", action="store_true", help="save training curves")
args = p.parse_args()
if args.epochs is None:
args.epochs = 20 if args.arch == "mlp" else 8
if args.out is None:
args.out = os.path.join(HERE, "checkpoints", f"{args.arch}.npz")
# np.savez appends .npz to a suffix-less path, so settle on the name it
# will actually write before anything derives the history/curve paths.
args.out = checkpoint_path(args.out)
return args
def accuracy(model, x, y, batch_size=1000):
return float((model.predict(x, batch_size) == y).mean())
def main():
args = parse_args()
rng = np.random.default_rng(args.seed)
d = data.load(val_size=args.val_size, seed=args.seed,
flatten=args.arch == "mlp")
config = {
"arch": args.arch,
"num_classes": d["num_classes"],
"activation": args.activation,
"batchnorm": args.batchnorm,
"dropout": args.dropout,
}
if args.arch == "mlp":
config.update(input_dim=int(d["input_dim"]), hidden=list(args.hidden))
else:
config.update(input_shape=d["input_shape"],
channels=list(args.channels),
kernel_size=args.kernel_size,
fc_hidden=list(args.fc_hidden))
model = build_model(config, rng=rng)
criterion = SoftmaxCrossEntropy()
if args.optimizer == "adam":
opt = Adam(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
else:
opt = SGD(model.parameters(), lr=args.lr, momentum=args.momentum,
nesterov=True, weight_decay=args.weight_decay)
# Batches reach the loop already normalised, so anything shifted in from
# outside the frame has to be filled with whatever 0 maps to — d["fill"].
# A literal 0.0 would paint a grey border around every warped digit.
augmenter = augmentation.build(d["image_shape"],
fill=d["fill"],
affine=args.augment,
elastic=args.elastic,
rotation=args.aug_rotation,
translate=args.aug_translate,
scale=args.aug_scale,
alpha=args.elastic_alpha,
sigma=args.elastic_sigma)
steps_per_epoch = int(np.ceil(len(d["x_train"]) / args.batch_size))
schedule = CosineSchedule(args.lr, args.epochs * steps_per_epoch,
warmup_steps=args.warmup, min_lr=args.lr * 0.01)
print(model)
print(f"{model.num_parameters():,} parameters | "
f"{len(d['x_train']):,} train / {len(d['x_val']):,} val / "
f"{len(d['x_test']):,} test")
print(f"{args.optimizer} lr={args.lr} batch={args.batch_size} "
f"epochs={args.epochs} dropout={args.dropout} bn={args.batchnorm}")
print(f"augment: {augmenter or 'off'}\n")
history = {"train_loss": [], "train_acc": [], "val_acc": [], "lr": []}
best_val, best_epoch = -1.0, -1
out_dir = os.path.dirname(args.out)
if out_dir: # a bare filename has no directory to make
os.makedirs(out_dir, exist_ok=True)
step = 0
t0 = time.perf_counter()
# BatchNorm needs two samples to have a variance at all, so a trailing
# batch of one is dropped rather than silently contributing nothing.
min_batch = 2 if args.batchnorm else 1
for epoch in range(1, args.epochs + 1):
running_loss, seen, correct = 0.0, 0, 0
for xb, yb in data.batches(d["x_train"], d["y_train"], args.batch_size,
rng, min_size=min_batch):
opt.lr = schedule(step)
if augmenter is not None:
# Train split only: val and test must stay untouched, or the
# reported accuracy is measured on a different distribution
# than the one anyone will actually use.
xb = augmenter(xb, rng)
logits = model.forward(xb, training=True)
loss = criterion.forward(logits, yb)
model.backward(criterion.backward())
opt.step()
running_loss += loss * len(xb)
correct += int((logits.argmax(axis=1) == yb).sum())
seen += len(xb)
step += 1
train_loss = running_loss / seen
train_acc = correct / seen
val_acc = accuracy(model, d["x_val"], d["y_val"]) if args.val_size else float("nan")
# Cast off NumPy scalars — np.float32 is not JSON serializable, and the
# history is written out at the end of the run.
history["train_loss"].append(float(train_loss))
history["train_acc"].append(float(train_acc))
history["val_acc"].append(float(val_acc))
history["lr"].append(float(opt.lr))
# With no validation split there is nothing to select on, so fall back
# to training accuracy. Selecting on val_acc there would compare
# against NaN, which loses every comparison — no checkpoint would ever
# be written and the reload below would report a *previous* run's
# weights as this run's result.
score = train_acc if not args.val_size else val_acc
marker = ""
if score > best_val:
best_val, best_epoch = score, epoch
save_checkpoint(args.out, model, config,
extra={"epoch": epoch, "val_acc": val_acc,
"mean": d["mean"], "std": d["std"],
"seed": args.seed, "val_size": args.val_size})
marker = " <- best, saved"
# flush: stdout is block-buffered when redirected to a file or pipe, so
# without this a long CNN run looks hung until it finishes.
print(f"epoch {epoch:>3}/{args.epochs} loss {train_loss:.4f} "
f"train {train_acc * 100:5.2f}% val {val_acc * 100:5.2f}% "
f"lr {opt.lr:.2e} {time.perf_counter() - t0:5.1f}s{marker}",
flush=True)
# Reload the best-validation weights before the single test measurement.
# Refusing to load when this run wrote nothing is the point: otherwise a
# zero-epoch or never-improving run silently measures whatever checkpoint
# happened to be sitting at that path and reports it as its own.
if best_epoch < 0:
raise SystemExit(f"no checkpoint was written (ran {args.epochs} "
f"epochs); refusing to report a stale result")
best_model, _, extra = load_checkpoint(args.out)
y_test = d["y_test"]
test_pred = best_model.predict(d["x_test"])
errors = int((test_pred != y_test).sum())
test_acc = float((test_pred == y_test).mean())
selected_on = "val" if args.val_size else "train"
print(f"\nbest {selected_on} {best_val * 100:.2f}% at epoch {best_epoch}")
print(f"test accuracy {test_acc * 100:.2f}% "
f"({errors} errors / {len(y_test)})")
print(f"checkpoint: {args.out}")
history_path = os.path.splitext(args.out)[0] + "_history.json"
with open(history_path, "w") as fh:
json.dump({"args": vars(args), "config": config, "history": history,
"best_val": best_val, "best_epoch": best_epoch,
"test_acc": test_acc}, fh, indent=2)
if args.plot:
plot_curves(history, os.path.splitext(args.out)[0] + "_curves.png")
def plot_curves(history, path):
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
epochs = range(1, len(history["train_loss"]) + 1)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4))
ax1.plot(epochs, history["train_loss"], color="#d62728")
ax1.set_xlabel("epoch")
ax1.set_ylabel("training loss")
ax1.grid(alpha=0.3)
ax2.plot(epochs, [a * 100 for a in history["train_acc"]], label="train")
ax2.plot(epochs, [a * 100 for a in history["val_acc"]], label="val")
ax2.set_xlabel("epoch")
ax2.set_ylabel("accuracy (%)")
ax2.legend()
ax2.grid(alpha=0.3)
fig.tight_layout()
fig.savefig(path, dpi=130)
print(f"curves: {path}")
if __name__ == "__main__":
main()