Building Models with nn.Module

Buffers vs parameters

Parameters are numbers the optimizer trains; buffers are numbers the model remembers — both are saved, both move to the GPU, and mixing them up corrupts checkpoints.

On this page 5
  1. Why buffers exist
  2. How saving fits in
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A parameter is a number the model learns; a buffer is a number the model writes down and keeps.

Think of a shopkeeper. Two kinds of knowledge run the shop. The first is skill — how to bargain, how to spot a fake note. That skill improves with practice, a little every day.

The second is the notebook — today's opening cash, the price list, who owes what. Nobody "practises" a notebook. It is state, recorded and looked up.

When the shop reopens tomorrow, both must survive the night. The skill and the notebook together are the shop. In PyTorch, the skill is the parameters and the notebook is the buffers.

Why buffers exist

Some numbers in a model must be remembered but must never be touched by training.

The clearest example: a model that scales its inputs using the average of the training data. That average was measured, not learned. If the optimizer nudged it, it would stop being the truth about the data.

Yet it must be saved with the model, and it must travel to the GPU with the model. So PyTorch gives such numbers their own drawer in the register — the buffer.

How saving fits in

             the model's register
        ┌──────────────┬──────────────┐
        │  parameters  │   buffers    │
        │  (trained)   │  (recorded)  │
        └──────┬───────┴──────┬───────┘
               └──────┬───────┘
                      v
               state_dict  →  saved file  →  restored model

The state dict is the model's complete written record — every parameter and every buffer, by name. Saving a model means writing that record to a file. Loading means reading it back into a model with the same shape.

A real example you have seen

Any AI feature that works offline on your phone was saved on a server and restored on the device. Its recorded statistics — the notebook — travelled inside the same file as its trained weights.

Remember this

  • Parameters are trained by the optimizer. Buffers are recorded by your code.
  • Both are saved in the state dict, and both move to the GPU with the model.
  • A tensor that is neither is not saved at all — that is the bug to fear.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU.

A module with both

normaliser.py
import torch
from torch import nn

class Normaliser(nn.Module):
    """Scales inputs using statistics measured once from training data."""
    def __init__(self):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(3))          # learned by the optimizer
        self.register_buffer("mean", torch.zeros(3))       # saved state, never trained
        self.register_buffer("std", torch.ones(3))

    def fit(self, data):
        self.mean.copy_(data.mean(dim=0))
        self.std.copy_(data.std(dim=0))

    def forward(self, x):
        return self.weight * (x - self.mean) / self.std

torch.manual_seed(0)
norm = Normaliser()
norm.fit(torch.tensor([[10., 200., 3.], [14., 260., 5.], [12., 230., 4.]]))

print("parameters:", [n for n, _ in norm.named_parameters()])
print("buffers:   ", [n for n, _ in norm.named_buffers()])
print("state_dict:", list(norm.state_dict().keys()))

torch.save(norm.state_dict(), "normaliser.pt")
fresh = Normaliser()
fresh.load_state_dict(torch.load("normaliser.pt", weights_only=True))
print("restored mean:", fresh.mean.tolist())
Output
parameters: ['weight']
buffers:    ['mean', 'std']
state_dict: ['weight', 'mean', 'std']
restored mean: [12.0, 230.0, 4.0]

The walkthrough

register_buffer("mean", ...) is a method call, not an assignment, because a buffer needs a name in the register and a plain tensor assignment would bypass the register entirely. After the call, self.mean works like any attribute.

copy_ writes new values into the existing buffer in place. Assigning self.mean = data.mean(dim=0) would replace the registered tensor with an unregistered one — the underscore method keeps the bookkeeping intact.

weights_only=True tells torch.load to accept tensors and nothing else. A checkpoint file can contain arbitrary pickled Python, and unpickling arbitrary Python runs it. Pass this flag whenever you load a file you did not write yourself. As of torch 2.6 it became the default; on 2.5 you say it explicitly.

Save the state dict, not the model. torch.save(model) pickles the whole object, including your class definition's import path. It breaks the moment you rename a file. The state dict is data only, and survives refactors.

The decision rule

The tensor is...Use
updated by gradient descentnn.Parameter
updated by your own code, must be savedregister_buffer
recomputed every forward, never saveda local variable in forward

The famous buffer in the wild is BatchNorm's running mean — measured during training, replayed at inference. That story gets its own lesson.

Common mistakes

Storing state as a plain attribute. self.mean = torch.zeros(3) works until you save. The state dict omits it, and the reloaded model scales with zeros. Nothing errors; predictions are quietly wrong.

Making measured state a Parameter. The optimizer will drag your measured statistics toward whatever lowers the loss. They stop being statistics.

Freezing instead of buffering. A parameter with requires_grad=False does work, but it lies about intent, and model.parameters() still yields it. If it is never trained by design, it is a buffer.

Loading across shape changes. load_state_dict demands matching names and shapes. After editing the architecture, load with strict=False and print the returned missing and unexpected keys — never ignore that return value.

Try it yourself

Add count — how many batches fit has seen — as a third buffer of shape (1,). Save, reload, and confirm the count survives the round trip.

What to learn next

Researcher — Mathematics and papers.

Registration and serialization semantics

register_buffer(name, tensor, persistent=True) inserts into Module._buffers. Persistent buffers appear in state_dict(); non-persistent ones (persistent=False) are registered — they move with .to() and appear in named_buffers() — but are excluded from serialization. Non-persistent buffers suit derived state that is cheap to rebuild, such as cached rotary-embedding tables in transformer implementations, where recomputation on load is preferable to bloating checkpoints.

Buffers differ from parameters in autograd only by type: a buffer is a plain Tensor, so it participates in the graph if it has requires_grad=True (unusual but legal), while nn.Parameter is a Tensor subclass that additionally self-registers on attribute assignment.

Checkpoint formats and their trade-offs

torch.save wraps pickle with a zip container of tensor storages. Its two structural weaknesses are documented: arbitrary-code execution on load (mitigated by weights_only=True, default since torch 2.6 — pin your expectation to the version you run) and no partial reads: loading one tensor materializes the whole archive's index. The safetensors format (Hugging Face) addresses both with a JSON header plus raw tensor bytes — no code execution by construction, and memory-mappable for zero-copy partial loads. For interchange, prefer it; for internal training checkpoints, torch.save with optimizer state remains the path of least resistance.

A full training checkpoint is more than the model:

{model.state_dict(), optimizer.state_dict(), scheduler.state_dict(),
 epoch, torch RNG state, numpy RNG state}

Omitting optimizer state silently resets Adam's first and second moments; resumed runs then repeat the warm-up transient, visible as a loss spike at the resume boundary.

Precision and device subtleties

state_dict() returns references, not copies — mutating a returned tensor mutates the live model. copy.deepcopy or {k: v.clone() for ...} when snapshotting mid-training. On load, load_state_dict copies into existing tensors, preserving device and (by default) dtype of the destination; loading fp32 checkpoints into an fp16 model quantizes at copy time without warning. torch.load(..., map_location="cpu") avoids accidentally requiring the GPU that saved the file.

Reading

  • The serialization notes in the PyTorch docs (docs.pytorch.org/docs/stable/notes/serialization.html) are the authoritative statement of forward/backward-compatibility promises.
  • Ioffe and Szegedy (2015), Batch Normalization — the paper that made running-statistics buffers ubiquitous.

What to learn next