Edge and On-device AI

Pruning

Pruning deletes the weights that contribute least, then retrains the survivors — but zeroed weights shrink nothing until something downstream exploits them.

On this page 8
  1. The short answer
  2. The analogy you have already lived
  3. Why it exists
  4. How it works
  5. The catch nobody mentions
  6. Two kinds of pruning
  7. Remember this
  8. 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.

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 further

Step 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

Developer — Code and libraries.

Setup

bash
pip install torch scikit-learn

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

prune_sweep.py
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))
Output
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.6

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

prune_structured.py
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)))
Output
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

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:

  1. 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.
  2. 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 typeAccuracy at a fixed parameter countWall-clock speed-up on commodity hardware
UnstructuredBestNone, unless a sparse kernel exists
2:4 semi-structuredSlightly worseUp to 2x on NVIDIA Ampere and later
Channel / headWorstProportional 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

What to learn next

What to learn next

These follow on from what you just read.

  • Edge and On-device AI

    Knowledge distillation

    Knowledge distillation trains a small model to copy a big model's full opinion rather than its final answer, which lets a 249x smaller network land within a point of its teacher.

  • Edge and On-device AI

    ONNX

    ONNX is a single open file format for trained models, so a model built in PyTorch can run in C++, Java, JavaScript or on a phone without shipping PyTorch with it.

  • Edge and On-device AI

    TensorFlow Lite

    TensorFlow Lite, now called LiteRT, converts a trained model into a small flat file that a few-megabyte runtime can execute on Android, iOS and microcontrollers.