After Training: Saving Models and Making Predictions
A trained network is just a few arrays of numbers, but using it later takes more than the arrays. Save it safely, resume training exactly where it stopped, feed it inputs prepared the same way as in training, and let it say when it isn't sure.
Inspect Architecture: Multi-layer perceptron1. The idea
At the end of Chapter 13 the digit reader lived only in memory. Close Python and it's gone, along with the minute of training it took; for a large model that would be days. So the first job is saving it. A trained network is nothing but its arrays, here W1, b1, W2 and b2, 101,770 numbers in all, and NumPy's .npz format stores named arrays in one file. It's also safe to open: an .npz file holds only numbers, while Python's general-purpose pickle format can run arbitrary code when a file is loaded, so a pickled 'model' from someone else is a program you are running.
The arrays alone aren't enough to use the model later, though. You also need the recipe that turned images into inputs (flatten, divide by 255), the architecture, and what each output means. And to resume training, you need more still: Adam's running averages and step counter (Chapter 10) and the shuffler's position in its random sequence. Saved together as a checkpoint, they let training continue exactly: two epochs, a save, a load into fresh objects and two more epochs give weights identical to four epochs without a break, to the last bit. Reload only the weights and the run drifts away, with weights differing by up to 0.15.
Using the model, called inference, needs only the forward pass: plain NumPy, no autograd graph. The important rule is that inputs must be prepared exactly as in training, and the two classic mistakes fail in different ways. Swap rows and columns and accuracy collapses to 14.8%, which you would notice. Forget to divide by 255 and accuracy barely changes, 96.8%, because a ReLU network is nearly indifferent to the scale of its input. But every score is now 255 times too large, so the probabilities are wrecked: 315 of the model's 322 mistakes come with a probability above 0.999.
Correct probabilities are worth protecting, because they let a model decline to answer. Accept only predictions with a top probability of at least 0.99 and the model answers two thirds of the test digits and gets 99.91% of those right, passing the rest to a person. Finally, the file can shrink: storing each weight as an 8-bit integer plus one scale per array cuts it from 814 KB to 102 KB, with the same test accuracy.
2. The math
to use the model: weights + input recipe + architecture + class names
to resume training: all that + optimizer state (Adam's m, v, t)
+ the shuffler's RNG state + the epochnp.savez / np.load: named arrays; np.load refuses pickled objects by default
pickle: rebuilds any Python object, and can run code doing it
JSON: text for the settings: epoch, recipe, versionsA: 4 epochs straight validation 96.88%
B: 2 epochs, save, load, 2 more validation 96.88%
weights identical to A (max diff 0.0)
C: same, but weights only reloaded validation 96.75%
weights differ from A by up to 0.153fresh Adam: t = 1, m̂ = g, v̂ = g²
⇒ step = η · sign(g) for every weight
with saved state: m and v hold 782 steps of history (2 epochs × 391)p = softmax( W₂ · ReLU(W₁ x + b₁) + b₂ )
plain NumPy: no graph, no gradients
saved file: 815,126 bytes (101,770 weights × 8 bytes + headers)
reloaded copy: identical probabilities on all 10,000 test images1,000 predictions one image at a time: tens of milliseconds
the same 1,000 as one (1000 × 784) batch: a few millisecondscorrect: 96.86%, mistakes with p > 0.999: 0
forgot to divide by 255: 96.78%, mistakes with p > 0.999: 315
rows and columns swapped: 14.84%, mistakes with p > 0.999: 152ReLU(c·u) = c·ReLU(u) for c > 0
biases small ⇒ scores(255·x) ≈ 255 · scores(x)
= softmax at temperature T = 1/255answer only if the top probability ≥ threshold
threshold 0: answers 100.0%, right on 96.86%
threshold 0.9: answers 89.0%, right on 99.45%
threshold 0.99: answers 67.1%, right on 99.91%
threshold 0.999: answers 33.6%, right on 100.00%float64: 814,160 bytes float32: 407,080 float16: 203,540
(test accuracy 96.86% for all three)
int8: q = round(w / s), s = max|w| / 127, restore w ≈ q · s
→ 101,770 bytes, test accuracy 96.87%
rounding error per weight ≤ s / 2• Adam's step for a weight is η · m̂ / (√v̂ + ε), where m and v are running averages of its past gradients and t counts steps (Chapter 10). The step depends on the whole history, not only on today's gradient.
• After 782 steps, m and v are settled averages. For a weight whose recent gradients point in different directions, m is small compared with √v, so its step is much smaller than η.
• Restart with m = v = 0 and t = 0. On the first step the bias correction gives m̂ = g and v̂ = g², so the step is η·g/|g| = η·sign(g): every weight moves by the full η, whether its gradient is large or tiny.
• That single jolt, plus a different order of mini-batches from the restarted shuffler, sends training down a different path. Nothing is wrong with it, but it isn't the run you paused, and it can't be repeated exactly. ∎
• Saving m, v, t and the shuffler's state removes every difference: the resumed run performs exactly the same arithmetic in the same order, which is why run B's weights match run A's to the last bit.
• ReLU is positively homogeneous: ReLU(c·u) = max(0, c·u) = c·max(0, u) = c·ReLU(u) for any c > 0. So is any matrix product: W(c·x) = c·Wx.
• Feed the network c·x instead of x, with c = 255. If the biases were zero, the hidden layer would be exactly c times larger, and so would every score: z(c·x) = c·z(x).
• The biases break this slightly, but they're small next to 255 times the weighted inputs, so z(c·x) ≈ c·z(x). Multiplying all scores by the same positive number never changes which is largest, so most predicted digits stay the same: 96.78% instead of 96.86%.
• Softmax of c·z is softmax of z at temperature T = 1/c (Chapter 9). With T = 1/255 the largest score takes essentially all the probability, so every prediction, right or wrong, comes out near 1. ∎
• Hence the dangerous combination: accuracy that looks normal and probabilities that are useless, 315 confident mistakes instead of 0. Checking a few predicted probabilities, not only accuracy, catches it.
3. How it works
4. The code (python)
# Chapter 14: after training. Save a model safely, resume training exactly, and use it to make predictions.
# Uses tensor.py (Chapter 6), optim.py (Chapter 10's Adam) and the mnist/ folder that Chapter 13 downloaded.
import gzip
import hashlib
import json
import os
import struct
import time
import numpy as np
from tensor import Tensor
from optim import Adam
# === 0. Chapter 13's data, read again ==============================================
def read_idx(name):
with gzip.open(os.path.join("mnist", name), "rb") as f:
data = f.read()
ndim = data[3]
shape = struct.unpack(">" + "I" * ndim, data[4:4 + 4 * ndim])
return np.frombuffer(data, dtype=np.uint8, offset=4 + 4 * ndim).reshape(shape)
images, labels = read_idx("train-images-idx3-ubyte.gz"), read_idx("train-labels-idx1-ubyte.gz")
test_images, test_labels = read_idx("t10k-images-idx3-ubyte.gz"), read_idx("t10k-labels-idx1-ubyte.gz")
def preprocess(raw):
"""The one place the input recipe lives: flatten to 784 numbers, scale to 0..1."""
return raw.reshape(len(raw), 784) / 255.0
X_train, y_train = preprocess(images[:50000]), labels[:50000]
X_val, y_val = preprocess(images[50000:]), labels[50000:]
X_test = preprocess(test_images)
# === 1. The model as named arrays, and a forward pass with no autograd at all =======
def new_model(seed):
rng = np.random.default_rng(seed)
return {"W1": rng.normal(0, np.sqrt(2 / 784), (128, 784)), "b1": np.zeros(128),
"W2": rng.normal(0, np.sqrt(2 / 128), (10, 128)), "b2": np.zeros(10)}
def predict_proba(model, X):
"""Inference only: plain NumPy, nothing recorded for a backward pass."""
h = np.maximum(0, X @ model["W1"].T + model["b1"])
z = h @ model["W2"].T + model["b2"]
e = np.exp(z - z.max(axis=1, keepdims=True))
return e / e.sum(axis=1, keepdims=True)
def accuracy(model, X, y):
return float(np.mean(predict_proba(model, X).argmax(axis=1) == y))
# === 2. Training that can stop and resume ============================================
def train(model, epochs, opt_state=None, rng_state=None, lr=0.001, batch=128):
"""Mini-batch Adam on the model's arrays. Returns everything needed to carry on later."""
params = {k: Tensor(v) for k, v in model.items()}
names = list(params)
opt = Adam([params[k] for k in names], lr)
if opt_state is not None: # put Adam's memory back
opt.t = opt_state["t"]
for i, k in enumerate(names):
opt.m[i][...] = opt_state["m_" + k]
opt.v[i][...] = opt_state["v_" + k]
shuffle = np.random.default_rng(0)
if rng_state is not None: # and the shuffler's place in its sequence
shuffle.bit_generator.state = rng_state
for _ in range(epochs):
order = shuffle.permutation(len(X_train))
for start in range(0, len(X_train), batch):
rows = order[start:start + batch]
h = (Tensor(X_train[rows]) @ params["W1"].T + params["b1"]).relu()
Z = h @ params["W2"].T + params["b2"]
E = (Z + (-Z.data.max(axis=1, keepdims=True))).exp()
P = E / (E @ Tensor(np.ones((10, 1))))
loss = -(Tensor(np.eye(10)[y_train[rows]]) * P.log()).sum() / len(rows)
for p in params.values():
p.grad = np.zeros_like(p.data)
loss.backward()
opt.step()
state = {"t": opt.t}
for i, k in enumerate(names):
state["m_" + k], state["v_" + k] = opt.m[i].copy(), opt.v[i].copy()
return {k: p.data for k, p in params.items()}, state, shuffle.bit_generator.state
def save_checkpoint(path, model, opt_state, rng_state, epoch):
"""Arrays in an .npz file (no code inside, unlike pickle); everything else as readable JSON."""
np.savez(path + ".npz", **model, **{"adam_" + k: v for k, v in opt_state.items() if k != "t"})
meta = {"epoch": epoch, "adam_t": opt_state["t"], "shuffle_state": rng_state,
"architecture": "784-128-10 ReLU, softmax", "input": "flatten 28x28, divide by 255",
"classes": list(range(10)), "numpy": np.__version__}
with open(path + ".json", "w") as f:
json.dump(meta, f, indent=1)
def load_checkpoint(path):
with np.load(path + ".npz") as arrays: # allow_pickle stays False: arrays only
model = {k: arrays[k] for k in ("W1", "b1", "W2", "b2")}
opt_state = {k[5:]: arrays[k] for k in arrays.files if k.startswith("adam_")}
with open(path + ".json") as f:
meta = json.load(f)
opt_state["t"] = meta["adam_t"]
return model, opt_state, meta
# Run A: four epochs without stopping
t0 = time.time()
model_a, _, _ = train(new_model(seed=0), epochs=4)
print(f"run A, 4 epochs straight: validation {accuracy(model_a, X_val, y_val):.2%} ({time.time() - t0:.0f} s)")
# Run B: two epochs, save, load into fresh objects, two more
model, opt_state, rng_state = train(new_model(seed=0), epochs=2)
save_checkpoint("checkpoint", model, opt_state, rng_state, epoch=2)
del model, opt_state, rng_state # as if the program had stopped
model, opt_state, meta = load_checkpoint("checkpoint")
model_b, _, _ = train(model, epochs=2, opt_state=opt_state, rng_state=meta["shuffle_state"])
diff_b = max(np.abs(model_a[k] - model_b[k]).max() for k in model_a)
print(f"run B, 2 + save/load + 2: validation {accuracy(model_b, X_val, y_val):.2%} largest weight difference from A: {diff_b}")
# Run C: the same, but forgetting Adam's state and the shuffler's position
model, _, meta = load_checkpoint("checkpoint")
model_c, _, _ = train(model, epochs=2)
diff_c = max(np.abs(model_a[k] - model_c[k]).max() for k in model_a)
print(f"run C, resumed with weights only: validation {accuracy(model_c, X_val, y_val):.2%} largest weight difference from A: {diff_c:.3f}")
# === 3. Making predictions with the saved model ===========================================
np.savez("digits_model.npz", **model_a)
with open("digits_model.npz", "rb") as f:
print(f"\nsaved model: {os.path.getsize('digits_model.npz'):,} bytes, sha256 {hashlib.sha256(f.read()).hexdigest()[:16]}…")
with np.load("digits_model.npz") as arrays:
served = {k: arrays[k] for k in arrays.files}
same = np.array_equal(predict_proba(served, X_test), predict_proba(model_a, X_test))
print(f"loaded copy gives identical probabilities on all 10,000 test images: {same}")
print(f"test accuracy {accuracy(served, X_test, test_labels):.2%}")
def predict_digit(model, raw_image):
"""One 28x28 image of 0..255 bytes in, (digit, probability) out, using the same preprocess as training."""
p = predict_proba(model, preprocess(raw_image[None]))[0]
return int(p.argmax()), float(p.max())
for i in range(3):
d, p = predict_digit(served, test_images[i])
print(f" test image {i}: label {test_labels[i]}, predicted {d} with probability {p:.4f}")
# Preprocessing that doesn't match training: two classic bugs
P_right = predict_proba(served, X_test)
for name, X_bad in [("forgot to divide by 255", test_images.reshape(-1, 784).astype(float)),
("rows and columns swapped", preprocess(test_images.transpose(0, 2, 1)))]:
P_bad = predict_proba(served, X_bad)
wrong = P_bad.argmax(axis=1) != test_labels
print(f"{name:25}: test accuracy {1 - wrong.mean():.2%}; mistakes made with probability > 0.999:"
f" {np.sum(wrong & (P_bad.max(axis=1) > 0.999))} (correct preprocessing: "
f"{np.sum((P_right.argmax(axis=1) != test_labels) & (P_right.max(axis=1) > 0.999))})")
# One at a time, or all together
t0 = time.time()
for i in range(1000):
predict_proba(served, X_test[i:i + 1])
one_by_one = time.time() - t0
t0 = time.time()
predict_proba(served, X_test[:1000])
batched = time.time() - t0
print(f"1,000 predictions one at a time: {one_by_one * 1000:.0f} ms; as one batch: {batched * 1000:.1f} ms")
# === 4. Knowing when not to answer =========================================================
P = predict_proba(served, X_test)
confidence, guess = P.max(axis=1), P.argmax(axis=1)
print("\nanswer only when the top probability is at least …")
for threshold in [0.0, 0.9, 0.99, 0.999]:
keep = confidence >= threshold
print(f" {threshold:5}: answers {keep.mean():6.1%} of test images, right on {np.mean(guess[keep] == test_labels[keep]):.2%} of those")
# === 5. Smaller files: fewer bits per weight =================================================
print("\nweights stored as …")
for name, convert in [("float64", lambda w: w),
("float32", lambda w: w.astype(np.float32)),
("float16", lambda w: w.astype(np.float16)),
("int8 + one scale per array", None)]:
if convert is None: # map each array's range onto −127..127
scales = {k: np.abs(w).max() / 127 for k, w in served.items()}
stored = {k: np.round(w / scales[k]).astype(np.int8) for k, w in served.items()}
restored = {k: stored[k].astype(float) * scales[k] for k in stored}
else:
stored = {k: convert(w) for k, w in served.items()}
restored = {k: v.astype(float) for k, v in stored.items()}
size = sum(v.nbytes for v in stored.values())
print(f" {name:27} {size:9,} bytes test accuracy {accuracy(restored, X_test, test_labels):.2%}")
5. Practice
Work these out on paper (or in Python) and type the number. Answers are checked with a small tolerance for rounding.
6. Go further
- Give each row of W1 its own int8 scale instead of one scale for the whole array (128 scales instead of 1). Compare the largest and average rounding error and the test accuracy with the per-array version. Then try 4 bits per weight (−7..7). Where does accuracy start to drop?
- Check whether the probabilities mean what they say: sort the test predictions into bins by top probability (0.5–0.6, 0.6–0.7, …, 0.99–1.0) and print each bin's actual accuracy. A well-calibrated model's 0.8 bin is right about 80% of the time. Is this one too confident, or not confident enough?
- Make training crash-proof: save a checkpoint after every epoch, then simulate a crash by stopping after epoch 7 of 10, start a new Python process that loads the latest checkpoint, and finish. Confirm the final weights equal an uninterrupted 10-epoch run with np.array_equal.
7. Check yourself
In the catalog
Finished this chapter?