Course Content
Deep Learning with TensorFlow and PyTorch
4 sections · 15 lessons
Saving and Loading Models
Fourteen hours into a training run, at epoch 47 of 60, the machine runs out of memory and the process dies. You have the terminal output. You do not have the model. Everything that mattered lived in GPU memory and is now gone.
That one is at least obvious. Here is the version that hurts more, because it looks like success.
You finish training, save the model, load it in a serving script, and validation accuracy has dropped from 94% to 71%. The weights are identical — you checked. What changed is that during training your inputs were standardised using a StandardScaler fitted on the training set, and your serving script feeds raw values. The model is fine. It is being handed data from a different universe than the one it learned in.
A model is not just its weights. It is the weights, the architecture, the preprocessing that produced its inputs, and the mapping from output index to label. Save three of the four and you have saved nothing usable.
Three different things "save the model" can mean
| Purpose | Must contain | Typical size |
|---|---|---|
| Checkpoint — resume an interrupted run | Weights, optimiser state, epoch, scheduler state, RNG state | 2–3× the weights (Adam holds two extra tensors per parameter) |
| Inference artefact — serve predictions | Weights, architecture, preprocessing, class names | About the size of the weights |
| Archive — reproduce this result later | All of the above plus code version, data version, config, metrics | Slightly larger; mostly metadata |
Conflating the first two is the usual mistake. A checkpoint with optimiser state is two to three times larger than it needs to be for serving, and an inference artefact without optimiser state cannot resume training properly — restarting Adam from scratch throws away the accumulated moment estimates and produces a visible bump in the loss curve.
Keras: what to call and what to avoid
The native format
1model.save("model.keras") # architecture + weights + optimiser state23import tensorflow as tf4loaded = tf.keras.models.load_model("model.keras")5loaded.predict(x) # ready immediately, no rebuild neededThe .keras extension is required — it selects the current zip-based format. It stores the architecture as configuration, the weights as arrays, and the optimiser state, so a reloaded model can continue training.
Weights only
1model.save_weights("weights.weights.h5")23fresh = build_model() # you must rebuild the same architecture4fresh.load_weights("weights.weights.h5")Use this when the architecture lives in code you control and you want the file to be just numbers — for transfer learning, or for loading pretrained weights into a modified model. The cost is that a file of weights alone is meaningless without the exact code that built the model.
Exporting for serving
1model.export("saved_model_dir") # a frozen inference graph, no Python2# Serve it without Keras:3# docker run -p 8501:8501 \4# -v "$(pwd)/saved_model_dir:/models/m/1" -e MODEL_NAME=m tensorflow/servingsave() and export() are for different jobs, and the distinction trips people up. save() produces a file you reload into Python to continue working with. export() produces a self-contained computation graph for a serving runtime — it cannot be trained further, and it does not need your model code to run. If you are shipping to TensorFlow Serving, LiteRT (formerly TF Lite) or TensorFlow.js, you want export().
The legacy .h5 whole-model format still loads for backwards compatibility, but it cannot represent everything a modern model can — custom layers and subclassed models in particular — and it should not be used for new work.
PyTorch: save the state_dict, not the object
PyTorch offers a tempting one-liner:
torch.save(model, "model.pt") # DON'Tmodel = torch.load("model.pt") # PyTorch 2.6+ refuses this unless weights_only=FalseIt pickles the Python object, which means the file contains a reference to your class by module path. Rename the file the class lives in, move it to a different package, or refactor the class, and the load fails with an import error. Worse, a pickle can execute arbitrary code on load, so a model file downloaded from the internet is an executable, not data.
1torch.save(model.state_dict(), "weights.pt") # DO23model = MyModel(hidden=256) # rebuild from code4model.load_state_dict(torch.load("weights.pt", weights_only=True))5model.eval() # essential before inferenceA state_dict is an ordered dictionary mapping parameter names to tensors. Nothing but names and numbers, so it survives refactoring and is safe to load.
weights_only=True restricts unpickling to plain tensors and refuses to execute arbitrary objects. Recent PyTorch versions default it to True, which occasionally breaks loading of older checkpoints that stored non-tensor objects — that error is the security feature working, and the correct response is to inspect the file's provenance rather than reflexively setting the flag back to False.
The model.eval() after loading is not optional. A freshly constructed module is in training mode, so dropout is active and batch norm uses batch statistics. Loading weights does not change the mode. Skip that line and every prediction is subtly wrong.
A full training checkpoint
1def save_checkpoint(path, model, opt, sched, epoch, best_val):2 torch.save({3 "epoch": epoch,4 "model_state": model.state_dict(),5 "optim_state": opt.state_dict(), # Adam's moment estimates6 "sched_state": sched.state_dict(), # where the LR schedule is7 "best_val": best_val,8 "rng_state": torch.get_rng_state(), # so data order continues identically9 }, path)1011def load_checkpoint(path, model, opt, sched):12 ckpt = torch.load(path, map_location="cpu", weights_only=True)13 model.load_state_dict(ckpt["model_state"])14 opt.load_state_dict(ckpt["optim_state"])15 sched.load_state_dict(ckpt["sched_state"])16 torch.set_rng_state(ckpt["rng_state"])17 return ckpt["epoch"] + 1, ckpt["best_val"] # resume from the NEXT epochmap_location="cpu" is a small detail with a large payoff. Without it, a checkpoint saved on GPU 3 of an eight-GPU machine tries to load onto GPU 3 specifically, and fails on any machine with fewer GPUs. Loading to CPU and then calling model.to(device) works everywhere.
Checkpointing during training
The rule is to save on every improvement, and to keep the last checkpoint as well so an interrupted run can resume from where it stopped rather than from the last time it improved.
1import os23best_val, wait, patience = float("inf"), 0, 104start_epoch = 05if os.path.exists("last.pt"): # resume if present6 start_epoch, best_val = load_checkpoint("last.pt", model, opt, sched)7 print(f"resuming from epoch {start_epoch}")89for epoch in range(start_epoch, num_epochs):10 train_one_epoch(model, opt, train_loader)11 val = evaluate(model, val_loader)12 sched.step()1314 save_checkpoint("last.pt", model, opt, sched, epoch, best_val) # always15 if val < best_val: # on improvement16 best_val, wait = val, 017 save_checkpoint("best.pt", model, opt, sched, epoch, best_val)18 else:19 wait += 120 if wait >= patience:21 break1tf.keras.callbacks.ModelCheckpoint(2 filepath="best.keras",3 monitor="val_loss",4 save_best_only=True, # otherwise you overwrite a good model with a worse one5 save_weights_only=False,6 verbose=1,7)save_best_only=True is doing real work. Without it, the final epoch's weights overwrite the best epoch's weights — and the final epoch is frequently past the point where the model started overfitting.
The other three-quarters of the artefact
This is where most deployment failures actually originate. Save the model alongside everything needed to reproduce its inputs and interpret its outputs.
1import json, joblib, hashlib, os, subprocess, torch2from datetime import datetime34def save_bundle(dirname, model, scaler, class_names, config, metrics):5 os.makedirs(dirname, exist_ok=True)67 torch.save(model.state_dict(), f"{dirname}/weights.pt")8 joblib.dump(scaler, f"{dirname}/scaler.joblib") # the exact fitted scaler910 with open(f"{dirname}/weights.pt", "rb") as f: # integrity check11 weights_sha256 = hashlib.sha256(f.read()).hexdigest()1213 meta = {14 "class_names": class_names, # index -> label, in the trained order15 "config": config, # every hyperparameter16 "metrics": metrics, # what it scored, so you can verify17 "input_shape": list(config["input_shape"]),18 "framework": torch.__version__,19 "created": datetime.now().isoformat(),20 "sha256": weights_sha256,21 "git_commit": subprocess.run(["git", "rev-parse", "HEAD"],22 capture_output=True, text=True).stdout.strip(),23 }24 with open(f"{dirname}/metadata.json", "w") as f:25 json.dump(meta, f, indent=2)The class_names entry prevents a genuinely nasty class of bug. If your training script derives labels from sorted(os.listdir(data_dir)) and someone adds a new category, every index shifts and your saved model starts confidently reporting the wrong class names — with unchanged accuracy metrics, because the weights are fine and only the interpretation is broken. Pin the ordering in the metadata and assert it at load time.
The single most valuable habit here: after saving, load the bundle in a fresh process and re-evaluate on a held-out batch. If the number does not match what training reported, something in the bundle is missing. Finding that now costs five minutes.
1def verify_bundle(dirname, X_check, y_check, expected_metric, tol=1e-4):2 model = build_model(**json.load(open(f"{dirname}/metadata.json"))["config"])3 model.load_state_dict(torch.load(f"{dirname}/weights.pt", weights_only=True))4 model.eval()5 scaler = joblib.load(f"{dirname}/scaler.joblib")6 with torch.no_grad():7 preds = model(torch.tensor(scaler.transform(X_check),8 dtype=torch.float32)).argmax(1).numpy()9 got = (preds == y_check).mean()10 assert abs(got - expected_metric) < tol, f"MISMATCH: {got} vs {expected_metric}"11 print("bundle verified:", got)Moving between frameworks: ONNX
ONNX is a common file format for computation graphs. Export from either framework and run in a runtime that has no dependency on PyTorch or TensorFlow at all — useful for C++ services, mobile, or a deployment environment where installing a full framework is impractical.
1import torch23dummy = torch.randn(1, 3, 224, 224)4torch.onnx.export( # uses the torch.export-based exporter5 model, (dummy,), "model.onnx",6 input_names=["input"], output_names=["logits"],7 dynamic_shapes=({0: torch.export.Dim("batch")},), # batch size may vary8)1import onnxruntime as ort2import numpy as np34sess = ort.InferenceSession("model.onnx", providers=["CPUExecutionProvider"])5out = sess.run(None, {"input": dummy.numpy()})[0]67# Always check the exported graph agrees with the original8np.testing.assert_allclose(out, model(dummy).detach().numpy(),9 rtol=1e-3, atol=1e-5)That final assertion is the point of the exercise. Since PyTorch 2.9, torch.onnx.export captures the model with torch.export by default (dynamo=True); the older dynamic_axes and low opset_version arguments still exist but now trigger warnings and conversions. The old TorchScript-based exporter traced Python control flow into a single fixed branch — a forward pass with if x.shape[0] > 1: exported whichever branch the dummy input took. The new exporter usually raises an error in that situation instead, which is better, but export can still change behaviour when an operation is translated approximately. Verify numerically, every time.
Making the artefact smaller
| Technique | Size reduction | Typical accuracy cost | When |
|---|---|---|---|
| float16 weights | 2× | Usually none measurable | Almost always safe for inference |
| Dynamic int8 quantisation | ~4× | 0–1% | CPU serving of dense and recurrent layers |
| Quantisation-aware training | ~4× | < 0.5% | When post-training quantisation costs too much accuracy |
| Magnitude pruning | 2–10× sparse | 0–2% | Only pays off with runtime support for sparsity |
| Distillation into a smaller model | 5–50× | 1–3% | When latency matters more than the last point of accuracy |
1# pip install torchao -- PyTorch's own torch.ao.quantization is deprecated2from torchao.quantization import quantize_, Int8DynamicActivationInt8WeightConfig34quantize_(model, Int8DynamicActivationInt8WeightConfig()) # Linear layers, in place5torch.save(model.state_dict(), "model_int8.pt")Measure accuracy after any of these, on the same held-out set, and record both numbers. "It got 4× smaller" is only half a result.
A save-and-load routine worth standardising on
Write the saving code at the same time as the training loop, not after the run has finished. The decisions are cheap to make early and expensive to retrofit.
Checkpoint continuously. Save last.pt every epoch and best.pt on every improvement. Any run longer than about twenty minutes should be resumable, because interruptions are not rare events — they are the normal way long runs end on shared machines and spot instances.
Save weights, not pickled objects. state_dict in PyTorch, .keras or weights files in TensorFlow. Your class definitions will change; the numbers will not.
Bundle the preprocessing with the model. The fitted scaler, the vocabulary, the image normalisation constants, the label ordering. This is the failure that opened this lesson and it remains the most common cause of "it worked in the notebook".
Record provenance. Git commit, config, dataset identifier, the metric the model achieved. Three months from now, someone — probably you — will need to know which of the eleven files in the models directory is the one that was deployed, and the only reliable answer is one written down at the time.
Verify in a fresh process. Load the artefact from disk in a new Python session, run it on a held-out batch, and assert the metric matches. This one check catches missing preprocessing, wrong mode, shifted class ordering, and broken exports — four separate failure modes, one assertion.