Pruning
Pruning deletes the weights that contribute least, then retrains the survivors — but zeroed weights shrink nothing until something downstream exploits them.
- 15 min read
- 3 reading levels
- Updated
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Pruning deletes the parts of a model that were barely doing anything, then retrains what is left.
The analogy you have already lived
Watch a gardener with a fruit tree. They cut away the thin, crossing, shaded branches. The tree looks smaller and slightly bare.
Next season it fruits better than before. The water and light now go to the branches that were already carrying weight.
A trained network is full of thin branches. Weights so close to zero that removing them changes almost nothing. Pruning cuts them, and then you let the tree grow back into its new shape.
Why it exists
A network needs to be big while it is learning. It explores many possibilities at once, and most of them lead nowhere.
At the end of training, a large share of the weights sit near zero. They were never wrong; they were never used. On a server that waste is invisible. On a phone you are paying storage, memory and battery for numbers that do nothing.
How it works
step 1 train the model normally
step 2 sort every weight by how large it is
step 3 set the smallest ones to zero
step 4 TRAIN AGAIN, so the survivors take over the work
step 5 repeat steps 2 to 4 if you want to go furtherStep 4 is the step people skip, and skipping it is why pruning gets a bad reputation. Cutting without retraining hurts badly. Cutting and then retraining often costs nothing at all.
The catch nobody mentions
Here is the part that surprises everyone. Setting a weight to zero does not make the file smaller.
A zero is still a number. It still takes the same four bytes on disk that any other number takes. Delete 80 out of every 100 weights and the file is exactly the same size.
You get the benefit only when something later takes advantage of the zeros:
- Zipping the file. Long runs of zeros compress beautifully. This works everywhere and needs nothing special.
- A sparse storage format that records only the non-zero weights and where they sit.
- Hardware built for it. Some chips can skip zeros during the arithmetic. Most cannot.
This is why the honest way to report pruning is to give the zipped size, not the raw one.
Two kinds of pruning
Unstructured pruning removes individual weights, scattered anywhere. It gives the best accuracy for a given number of weights removed. It usually delivers no speed-up, because the holes sit in random places.
Structured pruning removes whole neurons, whole channels or whole layers. It damages accuracy more for the same reduction. It also genuinely makes the model smaller and faster. What is left is an ordinary, smaller network.
On a phone, structured pruning is usually the one that pays.
Remember this
- Pruning zeroes the least useful weights, then retrains the rest.
- Zeroed weights do not shrink a file by themselves.
- Structured pruning removes whole neurons and actually speeds things up.
What to learn next
- Knowledge distillation — the fourth compression lever, and often the strongest.
- Latency and throughput — why a smaller model was not a faster one here.
- Overfitting and underfitting — why removing capacity sometimes helps.
Developer — Code and libraries.
Setup
pip install torch scikit-learnBoth scripts run on a CPU in about ten seconds and download nothing.
Part 1: how far can you cut, and does retraining save you?
torch.nn.utils.prune.l1_unstructured zeroes the smallest weights by absolute value. It works by attaching a mask; prune.remove then bakes the mask into the weights permanently.
import copy, gzip, os, torch, torch.nn as nn, torch.nn.functional as F
import torch.nn.utils.prune as prune
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
torch.manual_seed(0); torch.set_num_threads(1)
X, y = load_digits(return_X_y=True)
Xtr, Xte, ytr, yte = train_test_split(X / 16.0, y, test_size=0.3, random_state=0, stratify=y)
Xtr = torch.tensor(Xtr, dtype=torch.float32); ytr = torch.tensor(ytr, dtype=torch.long)
Xte = torch.tensor(Xte, dtype=torch.float32); yte = torch.tensor(yte, dtype=torch.long)
def fit(m, epochs, lr=1e-3):
opt = torch.optim.Adam(m.parameters(), lr=lr)
for _ in range(epochs):
for i in torch.randperm(len(Xtr)).split(64):
opt.zero_grad(); F.cross_entropy(m(Xtr[i]), ytr[i]).backward(); opt.step()
return m.eval()
def acc(m):
with torch.no_grad():
return (m(Xte).argmax(1) == yte).float().mean().item()
def sizes(m, path):
torch.save(m.state_dict(), path)
with open(path, "rb") as f, gzip.open(path + ".gz", "wb", 9) as g:
g.write(f.read()) # zipping is what turns zeros into saved bytes
return os.path.getsize(path) / 1024, os.path.getsize(path + ".gz") / 1024
model = fit(nn.Sequential(nn.Linear(64, 256), nn.ReLU(),
nn.Linear(256, 256), nn.ReLU(),
nn.Linear(256, 10)), 30)
kb, gz = sizes(model, "dense.pt")
print("sparsity accuracy accuracy after file KB zipped KB")
print(" retraining")
print(" 0%% %.4f - %7.1f %7.1f" % (acc(model), kb, gz))
for amount in [0.5, 0.7, 0.8, 0.9, 0.95]:
m = copy.deepcopy(model)
lins = [l for l in m if isinstance(l, nn.Linear)]
for l in lins:
prune.l1_unstructured(l, name="weight", amount=amount)
before = acc(m) # damage, before any repair
torch.manual_seed(1); fit(m, 15, lr=5e-4) # the survivors take over the work
for l in lins:
prune.remove(l, "weight") # bake the mask in, drop the bookkeeping
kb, gz = sizes(m, "p.pt")
print("%6.0f%% %.4f %.4f %7.1f %7.1f" % (amount * 100, before, acc(m), kb, gz))sparsity accuracy accuracy after file KB zipped KB
retraining
0% 0.9648 - 334.2 308.6
50% 0.9648 0.9778 334.1 185.9
70% 0.9574 0.9815 334.1 126.9
80% 0.8593 0.9778 334.1 95.3
90% 0.6019 0.9537 334.1 61.5
95% 0.2426 0.8815 334.1 42.6Four things in that table
Retraining is the entire technique. At 80% sparsity, cutting alone drops accuracy to 0.8593. Fifteen epochs of retraining brings it to 0.9778 — above the original 0.9648. At 95% the difference is 0.2426 against 0.8815. Anyone reporting pruning damage without retraining is measuring the wrong thing.
Pruning can improve a small model. 0.9815 at 70% sparsity, against 0.9648 dense. Removing capacity acts as a regulariser here, in the same way described in overfitting and underfitting. On a large dataset with a well-fitted model, do not expect this. Here the dataset is 1257 images and the model has 85,002 parameters, so there was capacity to spare.
The file size column never moves. 334.1 KB at every sparsity level. This is the point from the beginner block, shown rather than claimed. Zeros cost four bytes each.
The zipped column is the real one. 308.6 KB down to 42.6 KB. That is a 7.2x reduction in what a user downloads, and it needed no special format or hardware — gzip did all of it.
Part 2: structured pruning, which actually makes the model smaller
Unstructured pruning leaves a full-size matrix full of holes. Structured pruning removes whole neurons, so what remains is an ordinary, genuinely smaller network.
import os, time, torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
torch.manual_seed(0); torch.set_num_threads(1)
X, y = load_digits(return_X_y=True)
Xtr, Xte, ytr, yte = train_test_split(X / 16.0, y, test_size=0.3, random_state=0, stratify=y)
Xtr = torch.tensor(Xtr, dtype=torch.float32); ytr = torch.tensor(ytr, dtype=torch.long)
Xte = torch.tensor(Xte, dtype=torch.float32); yte = torch.tensor(yte, dtype=torch.long)
def fit(m, epochs, lr=1e-3):
opt = torch.optim.Adam(m.parameters(), lr=lr)
for _ in range(epochs):
for i in torch.randperm(len(Xtr)).split(64):
opt.zero_grad(); F.cross_entropy(m(Xtr[i]), ytr[i]).backward(); opt.step()
return m.eval()
def acc(m):
with torch.no_grad(): return (m(Xte).argmax(1) == yte).float().mean().item()
def kb(m, p):
torch.save(m.state_dict(), p); return os.path.getsize(p) / 1024
def lat(m, n=400):
one = Xte[:1]
with torch.no_grad():
for _ in range(40): m(one)
t = time.perf_counter()
for _ in range(n): m(one)
return (time.perf_counter() - t) / n * 1000
model = fit(nn.Sequential(nn.Linear(64, 256), nn.ReLU(),
nn.Linear(256, 256), nn.ReLU(),
nn.Linear(256, 10)), 30)
KEEP = 64
l0, l2, l4 = model[0], model[2], model[4]
keep1 = l0.weight.abs().sum(1).topk(KEEP).indices.sort().values # busiest neurons in layer 1
keep2 = l2.weight.abs().sum(1).topk(KEEP).indices.sort().values # busiest neurons in layer 2
small = nn.Sequential(nn.Linear(64, KEEP), nn.ReLU(),
nn.Linear(KEEP, KEEP), nn.ReLU(),
nn.Linear(KEEP, 10))
with torch.no_grad():
small[0].weight.copy_(l0.weight[keep1]); small[0].bias.copy_(l0.bias[keep1])
small[2].weight.copy_(l2.weight[keep2][:, keep1]); small[2].bias.copy_(l2.bias[keep2])
small[4].weight.copy_(l4.weight[:, keep2]); small[4].bias.copy_(l4.bias)
small.eval()
print("model params file KB accuracy ms/image")
print("dense, 256 wide %8d %8.1f %.4f %.3f"
% (sum(p.numel() for p in model.parameters()), kb(model, "a.pt"), acc(model), lat(model)))
print("64 strongest neurons kept %8d %8.1f %.4f %.3f"
% (sum(p.numel() for p in small.parameters()), kb(small, "b.pt"), acc(small), lat(small)))
torch.manual_seed(1); fit(small, 15, lr=5e-4)
print("the same, after retraining %8d %8.1f %.4f %.3f"
% (sum(p.numel() for p in small.parameters()), kb(small, "c.pt"), acc(small), lat(small)))model params file KB accuracy ms/image dense, 256 wide 85002 334.1 0.9648 0.022 64 strongest neurons kept 8970 37.1 0.5870 0.021 the same, after retraining 8970 37.1 0.9630 0.022
The parameter counts, file sizes and accuracies are exact and will reproduce. The ms/image column is a measurement on one machine and will differ on yours.
What structured pruning bought, and what it did not
The file did shrink: 334.1 KB to 37.1 KB, 9 times smaller, with 9.5 times fewer parameters. That is the difference from Part 1, where the file never moved.
Accuracy came back: 0.5870 after the cut, 0.9630 after retraining. Almost the dense model's 0.9648, from a network with a tenth of the parameters.
Latency did not budge. About 0.02 ms in every row, and the small differences between rows are measurement noise. A ten-times smaller model that is exactly as fast looks broken until you think about what is being measured. At this size the time goes on Python call overhead and framework dispatch, not on the arithmetic. Shrinking the arithmetic changed nothing, because the arithmetic was never the bottleneck.
That is not a flaw in structured pruning. It is a lesson about measurement: shrink the thing that is actually costing you time, and read latency and throughput before optimising anything.
Common mistakes
Pruning without retraining. The single most common error. Compare the two accuracy columns in Part 1 again.
Cutting everything to the same sparsity. The first and last layers are far more sensitive than the middle ones, and the last layer is usually small enough that pruning it buys nothing. Production recipes prune per layer, with the input and output layers exempted or lightly pruned.
Going straight to the target sparsity. Cutting from 0% to 95% in one step damages more than five steps of 19% with retraining in between. Iterative pruning is standard for a reason.
Reporting the unzipped size. It makes unstructured pruning look useless, which is the mistake this lesson exists to prevent.
Expecting a speed-up from unstructured sparsity. Dense kernels multiply by zero at full speed. Without sparse kernels or sparsity-aware hardware, you have saved bytes and nothing else.
Try it yourself
In Part 1, delete the fit(m, 15, lr=5e-4) line and rerun. The "after retraining" column collapses to match the "before" column. Then put it back and change 15 epochs to 40. You will find a point where more retraining stops helping, and that point is your real sparsity budget.
What to learn next
- Knowledge distillation — the fourth compression lever, and often the strongest.
- Latency and throughput — why a smaller model was not a faster one here.
- Overfitting and underfitting — why removing capacity sometimes helps.
Researcher — Mathematics and papers.
The problem statement
Pruning seeks a mask $m \in {0,1}^n$ applied to parameters $\theta \in \mathbb{R}^n$:
$$ \min_{m, \theta} \; \mathcal{L}(f_{m \odot \theta}) \quad \text{s.t.} \quad |m|_0 \le k $$
- $\odot$ — elementwise product, so $m_i = 0$ deletes parameter $i$.
- $|m|_0$ — the number of surviving parameters; $k$ is the budget.
- $\mathcal{L}$ — loss on the deployment distribution, not the training set.
This is combinatorial and NP-hard in general. Every practical method is a heuristic for choosing $m$, followed by gradient descent on the surviving $\theta$.
Saliency criteria
Magnitude. Score $s_i = |\theta_i|$. Extremely cheap, and repeatedly competitive. Han et al. (2015), Learning both Weights and Connections, established it as the baseline everything is measured against.
Second-order. LeCun et al. (1989), Optimal Brain Damage, expand the loss around a trained minimum. With the gradient near zero, the change from setting $\theta_i \to 0$ is
$$ \delta \mathcal{L}i \approx \tfrac{1}{2} H{ii}\, \theta_i^2 $$
where $H_{ii}$ is the $i$-th diagonal element of the Hessian. Optimal Brain Surgeon (Hassibi and Stork, 1993) drops the diagonal assumption and additionally computes a compensating update to the remaining weights:
$$ \delta \mathcal{L}_i = \frac{\theta_i^2}{2\,[H^{-1}]_{ii}}, \qquad \delta \theta = -\frac{\theta_i}{[H^{-1}]_{ii}} H^{-1} e_i $$
This is exactly the machinery GPTQ reuses for quantisation, and SparseGPT (Frantar and Alistarh, 2023) reuses for one-shot pruning of large language models — layer-wise, so the $O(d^3)$ inverse stays tractable.
Movement. Sanh et al. (2020) score by how far a weight moves away from zero during fine-tuning, rather than by its final magnitude. This is the right criterion for transfer learning, where a large pretrained weight may be irrelevant to the downstream task.
The lottery ticket hypothesis
Frankle and Carbin (2019) showed that a dense network contains a sparse subnetwork which, reset to its original initialisation and trained in isolation, matches the dense network's accuracy in comparable time. The mask is found by training, pruning by magnitude, and rewinding.
Two important qualifications, both established by follow-up work rather than the original paper:
- Rewinding to initialisation fails at scale. Frankle et al. (2020) found that rewinding to an early training iteration rather than to step zero is required for ResNet-50 and larger.
- The winning masks are not reliably transferable, and the search costs more compute than training the dense model once.
The hypothesis is scientifically important and operationally weak. Nobody prunes production models this way.
Structured versus unstructured, quantified
The gap is entirely about hardware, not statistics.
| Sparsity type | Accuracy at a fixed parameter count | Wall-clock speed-up on commodity hardware |
|---|---|---|
| Unstructured | Best | None, unless a sparse kernel exists |
| 2:4 semi-structured | Slightly worse | Up to 2x on NVIDIA Ampere and later |
| Channel / head | Worst | Proportional to the reduction, on any hardware |
NVIDIA's 2:4 pattern — exactly two non-zeros in every group of four contiguous weights — is the compromise that made sparsity real on GPUs. Mishra et al. (2021) report near-baseline accuracy across a range of networks with the mandated retraining recipe.
Gale, Elsen and Hooker (2019), The State of Sparsity in Deep Neural Networks, is the honest audit: after controlling for training budget, simple magnitude pruning matched or beat the more elaborate methods on ResNet-50 and Transformer.
Sparsity and fairness
Hooker et al. (2020), What Do Compressed Deep Neural Networks Forget?, found that pruning does not distribute its damage evenly. Aggregate accuracy holds while error concentrates on a small subset of inputs — disproportionately underrepresented classes and atypical examples. They call these pruning-identified exemplars.
The operational consequence is direct: a compression report showing "0.2% top-1 loss" can conceal a large regression on a minority subgroup. Evaluate compressed models per-slice, not only in aggregate.
Reading
- Han et al., Learning both Weights and Connections for Efficient Neural Networks, NeurIPS 2015 — arxiv.org/abs/1506.02626
- LeCun, Denker and Solla, Optimal Brain Damage, NeurIPS 1989.
- Frankle and Carbin, The Lottery Ticket Hypothesis, ICLR 2019 — arxiv.org/abs/1803.03635
- Gale, Elsen and Hooker, The State of Sparsity in Deep Neural Networks, 2019 — arxiv.org/abs/1902.09574
- Hooker et al., What Do Compressed Deep Neural Networks Forget?, 2020 — arxiv.org/abs/1911.05248
- Frantar and Alistarh, SparseGPT, ICML 2023 — arxiv.org/abs/2301.00774
What to learn next
- Knowledge distillation — the fourth compression lever, and often the strongest.
- Latency and throughput — why a smaller model was not a faster one here.
- Overfitting and underfitting — why removing capacity sometimes helps.