Deep Learning with TensorFlow and PyTorch

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.

What has to be in the artefact, not just the weightsstate_dict ofthe weightsArchitectureor class codePreprocessingand scaler statsClass indexto label mapFramework andversion notetopbottomWeights alone reload into a model that scales its inputs differently and answers wrongly.
The weights are the quarter that people save and the three quarters that follow are what make them usable.

Three different things "save the model" can mean

PurposeMust containTypical size
Checkpoint — resume an interrupted runWeights, optimiser state, epoch, scheduler state, RNG state2–3× the weights (Adam holds two extra tensors per parameter)
Inference artefact — serve predictionsWeights, architecture, preprocessing, class namesAbout the size of the weights
Archive — reproduce this result laterAll of the above plus code version, data version, config, metricsSlightly 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

Python
model.save("model.keras")                 # architecture + weights + optimiser stateimport tensorflow as tfloaded = tf.keras.models.load_model("model.keras")loaded.predict(x)                          # ready immediately, no rebuild needed

The .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

Python
model.save_weights("weights.weights.h5")fresh = build_model()                      # you must rebuild the same architecturefresh.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

Python
model.export("saved_model_dir")            # a frozen inference graph, no Python# Serve it without Keras:#   docker run -p 8501:8501 \#     -v "$(pwd)/saved_model_dir:/models/m/1" -e MODEL_NAME=m tensorflow/serving

save() 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:

Python
torch.save(model, "model.pt")             # DON'Tmodel = torch.load("model.pt")            # PyTorch 2.6+ refuses this unless weights_only=False

It 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.

Python
torch.save(model.state_dict(), "weights.pt")          # DOmodel = MyModel(hidden=256)                            # rebuild from codemodel.load_state_dict(torch.load("weights.pt", weights_only=True))model.eval()                                           # essential before inference

A 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

Python
def save_checkpoint(path, model, opt, sched, epoch, best_val):    torch.save({        "epoch": epoch,        "model_state": model.state_dict(),        "optim_state": opt.state_dict(),      # Adam's moment estimates        "sched_state": sched.state_dict(),    # where the LR schedule is        "best_val": best_val,        "rng_state": torch.get_rng_state(),   # so data order continues identically    }, path)def load_checkpoint(path, model, opt, sched):    ckpt = torch.load(path, map_location="cpu", weights_only=True)    model.load_state_dict(ckpt["model_state"])    opt.load_state_dict(ckpt["optim_state"])    sched.load_state_dict(ckpt["sched_state"])    torch.set_rng_state(ckpt["rng_state"])    return ckpt["epoch"] + 1, ckpt["best_val"]   # resume from the NEXT epoch

map_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.

Python
import osbest_val, wait, patience = float("inf"), 0, 10start_epoch = 0if os.path.exists("last.pt"):                                 # resume if present    start_epoch, best_val = load_checkpoint("last.pt", model, opt, sched)    print(f"resuming from epoch {start_epoch}")for epoch in range(start_epoch, num_epochs):    train_one_epoch(model, opt, train_loader)    val = evaluate(model, val_loader)    sched.step()    save_checkpoint("last.pt", model, opt, sched, epoch, best_val)   # always    if val < best_val:                                               # on improvement        best_val, wait = val, 0        save_checkpoint("best.pt", model, opt, sched, epoch, best_val)    else:        wait += 1        if wait >= patience:            break
Python
tf.keras.callbacks.ModelCheckpoint(    filepath="best.keras",    monitor="val_loss",    save_best_only=True,     # otherwise you overwrite a good model with a worse one    save_weights_only=False,    verbose=1,)

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.

Python
import json, joblib, hashlib, os, subprocess, torchfrom datetime import datetimedef save_bundle(dirname, model, scaler, class_names, config, metrics):    os.makedirs(dirname, exist_ok=True)    torch.save(model.state_dict(), f"{dirname}/weights.pt")    joblib.dump(scaler, f"{dirname}/scaler.joblib")          # the exact fitted scaler    with open(f"{dirname}/weights.pt", "rb") as f:           # integrity check        weights_sha256 = hashlib.sha256(f.read()).hexdigest()    meta = {        "class_names": class_names,        # index -> label, in the trained order        "config": config,                  # every hyperparameter        "metrics": metrics,                # what it scored, so you can verify        "input_shape": list(config["input_shape"]),        "framework": torch.__version__,        "created": datetime.now().isoformat(),        "sha256": weights_sha256,        "git_commit": subprocess.run(["git", "rev-parse", "HEAD"],                                     capture_output=True, text=True).stdout.strip(),    }    with open(f"{dirname}/metadata.json", "w") as f:        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.

Python
def verify_bundle(dirname, X_check, y_check, expected_metric, tol=1e-4):    model = build_model(**json.load(open(f"{dirname}/metadata.json"))["config"])    model.load_state_dict(torch.load(f"{dirname}/weights.pt", weights_only=True))    model.eval()    scaler = joblib.load(f"{dirname}/scaler.joblib")    with torch.no_grad():        preds = model(torch.tensor(scaler.transform(X_check),                                   dtype=torch.float32)).argmax(1).numpy()    got = (preds == y_check).mean()    assert abs(got - expected_metric) < tol, f"MISMATCH: {got} vs {expected_metric}"    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.

Python
import torchdummy = torch.randn(1, 3, 224, 224)torch.onnx.export(                         # uses the torch.export-based exporter    model, (dummy,), "model.onnx",    input_names=["input"], output_names=["logits"],    dynamic_shapes=({0: torch.export.Dim("batch")},),   # batch size may vary)
Python
import onnxruntime as ortimport numpy as npsess = ort.InferenceSession("model.onnx", providers=["CPUExecutionProvider"])out = sess.run(None, {"input": dummy.numpy()})[0]# Always check the exported graph agrees with the originalnp.testing.assert_allclose(out, model(dummy).detach().numpy(),                           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

TechniqueSize reductionTypical accuracy costWhen
float16 weights2×Usually none measurableAlmost 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 pruning2–10× sparse0–2%Only pays off with runtime support for sparsity
Distillation into a smaller model5–50×1–3%When latency matters more than the last point of accuracy
Python
# pip install torchao -- PyTorch's own torch.ao.quantization is deprecatedfrom torchao.quantization import quantize_, Int8DynamicActivationInt8WeightConfigquantize_(model, Int8DynamicActivationInt8WeightConfig())   # Linear layers, in placetorch.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.